mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
compile backward
This commit is contained in:
1 parent
fefdea2304
commit
69489225c3
6 files changed
+216
-27
No files matched your search
@@ -12,3 +12,4 @@ examples/tinystories-llm/TinyStoriesV2-GPT4-valid.txt
|
||||
|
||||
# Environment files
|
||||
.env
|
||||
AGENTS.md
|
||||
@@ -18,7 +18,7 @@ loss.backward()
|
||||
optimizer.step()
|
||||
```
|
||||
|
||||
MLX implements this sequence by recording the outer model call and replaying it inside `value_and_grad` during `backward`. Random state is restored for the replay so dropout uses the same mask. Unrecorded loss expressions, gradient hooks, parameter `.grad`, higher-order gradients, and multiple-forward losses remain unsupported.
|
||||
MLX implements this sequence by recording the outer model call and replaying it inside a cached, compiled `value_and_grad` during `backward`. The optimizer update is compiled separately. Inputs and loss operands remain dynamic, so batches are not captured as constants. Random state is restored for the replay so dropout uses the same mask. Models with eager data-dependent operations fall back to an uncompiled replay. Unrecorded loss expressions, gradient hooks, parameter `.grad`, higher-order gradients, and multiple-forward losses remain unsupported.
|
||||
|
||||
Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation.
|
||||
|
||||
|
||||
+139
-15
@@ -11,6 +11,78 @@ _replayed_losses = {}
|
||||
_suspended = False
|
||||
|
||||
|
||||
class _ArraySlot:
|
||||
def __init__(self, index):
|
||||
self.index = index
|
||||
|
||||
|
||||
class _BackwardPlan:
|
||||
def __init__(self, value_and_grad, state):
|
||||
self._value_and_grad = value_and_grad
|
||||
self._compiled = mx.compile(
|
||||
value_and_grad,
|
||||
inputs=state,
|
||||
outputs=state,
|
||||
)
|
||||
self._compile_enabled = True
|
||||
|
||||
def __call__(self, *inputs):
|
||||
function = self._compiled if self._compile_enabled else self._value_and_grad
|
||||
return function(*inputs)
|
||||
|
||||
def fallback(self, *inputs):
|
||||
self._compile_enabled = False
|
||||
return self._value_and_grad(*inputs)
|
||||
|
||||
|
||||
def _partition_arrays(value, arrays):
|
||||
if isinstance(value, mx.array):
|
||||
slot = _ArraySlot(len(arrays))
|
||||
arrays.append(value)
|
||||
return slot
|
||||
if isinstance(value, tuple):
|
||||
return tuple(_partition_arrays(item, arrays) for item in value)
|
||||
if isinstance(value, list):
|
||||
return [_partition_arrays(item, arrays) for item in value]
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
key: _partition_arrays(item, arrays) for key, item in value.items()
|
||||
}
|
||||
return value
|
||||
|
||||
|
||||
def _restore_arrays(template, arrays):
|
||||
if isinstance(template, _ArraySlot):
|
||||
return arrays[template.index]
|
||||
if isinstance(template, tuple):
|
||||
return tuple(_restore_arrays(item, arrays) for item in template)
|
||||
if isinstance(template, list):
|
||||
return [_restore_arrays(item, arrays) for item in template]
|
||||
if isinstance(template, dict):
|
||||
return {
|
||||
key: _restore_arrays(item, arrays) for key, item in template.items()
|
||||
}
|
||||
return template
|
||||
|
||||
|
||||
def _template_key(template):
|
||||
if isinstance(template, _ArraySlot):
|
||||
return ("array",)
|
||||
if isinstance(template, tuple):
|
||||
return ("tuple", tuple(_template_key(item) for item in template))
|
||||
if isinstance(template, list):
|
||||
return ("list", tuple(_template_key(item) for item in template))
|
||||
if isinstance(template, dict):
|
||||
return (
|
||||
"dict",
|
||||
tuple(
|
||||
(type(key).__qualname__, repr(key), _template_key(item))
|
||||
for key, item in template.items()
|
||||
),
|
||||
)
|
||||
return ("static", type(template).__qualname__, repr(template))
|
||||
|
||||
|
||||
def _snapshot_random_state():
|
||||
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
|
||||
mx.eval(state)
|
||||
@@ -38,7 +110,7 @@ def end_forward(model, args, kwargs, output):
|
||||
and not _suspended
|
||||
and isinstance(output, mx.array)
|
||||
):
|
||||
_latest_forward = (model, args, kwargs, lambda value: value)
|
||||
_latest_forward = (model, args, kwargs, ())
|
||||
_lineage[id(output)] = (output, _latest_forward)
|
||||
|
||||
|
||||
@@ -57,37 +129,69 @@ def activate(optimizer):
|
||||
optimizer._random_before_forward = _snapshot_random_state()
|
||||
|
||||
|
||||
def propagate(source, result, operation=lambda value: value):
|
||||
def propagate(source, result, operation, signature, operands=()):
|
||||
entry = _lineage.get(id(source))
|
||||
if not _suspended and entry is not None and entry[0] is source:
|
||||
model, args, kwargs, previous = entry[1]
|
||||
model, args, kwargs, transforms = entry[1]
|
||||
_lineage[id(result)] = (
|
||||
result,
|
||||
(
|
||||
model,
|
||||
args,
|
||||
kwargs,
|
||||
lambda output: operation(previous(output)),
|
||||
transforms + ((operation, signature, operands),),
|
||||
),
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def register_loss(loss, loss_input, rebuild):
|
||||
def register_loss(loss, loss_input, rebuild, operands, signature):
|
||||
if _suspended or _active_optimizer is None:
|
||||
return loss
|
||||
entry = _lineage.get(id(loss_input))
|
||||
if entry is None or entry[0] is not loss_input:
|
||||
return loss
|
||||
model, args, kwargs, transform = entry[1]
|
||||
model, args, kwargs, transforms = entry[1]
|
||||
|
||||
def plan():
|
||||
def objective():
|
||||
return rebuild(transform(model(*args, **kwargs)))
|
||||
context = (
|
||||
args,
|
||||
kwargs,
|
||||
tuple(transform_operands for _, _, transform_operands in transforms),
|
||||
operands,
|
||||
)
|
||||
dynamic_inputs = []
|
||||
template = _partition_arrays(context, dynamic_inputs)
|
||||
cache_key = (
|
||||
getattr(model, "training", None),
|
||||
_template_key(template),
|
||||
tuple(transform_signature for _, transform_signature, _ in transforms),
|
||||
signature,
|
||||
)
|
||||
|
||||
return nn.value_and_grad(model, objective)()
|
||||
def build():
|
||||
state = [model.state, mx.random.state]
|
||||
|
||||
_loss_plans[id(loss)] = (loss, model, plan)
|
||||
def objective(*current_inputs):
|
||||
current_args, current_kwargs, current_transform_operands, current_operands = (
|
||||
_restore_arrays(template, current_inputs)
|
||||
)
|
||||
output = model(*current_args, **current_kwargs)
|
||||
for (operation, _, _), operation_operands in zip(
|
||||
transforms, current_transform_operands
|
||||
):
|
||||
output = operation(output, *operation_operands)
|
||||
return rebuild(output, *current_operands)
|
||||
|
||||
value_and_grad = nn.value_and_grad(model, objective)
|
||||
return _BackwardPlan(value_and_grad, state)
|
||||
|
||||
_loss_plans[id(loss)] = (
|
||||
loss,
|
||||
model,
|
||||
cache_key,
|
||||
tuple(dynamic_inputs),
|
||||
build,
|
||||
)
|
||||
return loss
|
||||
|
||||
|
||||
@@ -100,12 +204,32 @@ def backward(loss):
|
||||
raise RuntimeError(
|
||||
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
|
||||
)
|
||||
_, model, plan = entry
|
||||
_, model, cache_key, dynamic_inputs, build = entry
|
||||
random_after_forward = _snapshot_random_state()
|
||||
_restore_random_state(_active_optimizer._random_before_forward)
|
||||
_suspended = True
|
||||
try:
|
||||
replayed_loss, gradients = plan()
|
||||
plan = _active_optimizer._compiled_backward.get(cache_key)
|
||||
if plan is None:
|
||||
plan = build()
|
||||
_active_optimizer._compiled_backward[cache_key] = plan
|
||||
parameters_before_replay = model.trainable_parameters()
|
||||
try:
|
||||
replayed_loss, gradients = plan(*dynamic_inputs)
|
||||
except ValueError as error:
|
||||
message = str(error)
|
||||
if (
|
||||
not plan._compile_enabled
|
||||
or "Attempting to eval an array" not in message
|
||||
or not (
|
||||
"function transformations" in message
|
||||
or "without a primitive" in message
|
||||
)
|
||||
):
|
||||
raise
|
||||
model.update(parameters_before_replay)
|
||||
_restore_random_state(_active_optimizer._random_before_forward)
|
||||
replayed_loss, gradients = plan.fallback(*dynamic_inputs)
|
||||
mx.eval(replayed_loss, gradients, mx.random.state)
|
||||
finally:
|
||||
_suspended = False
|
||||
@@ -126,8 +250,8 @@ def step(optimizer):
|
||||
if optimizer._pending_update is None:
|
||||
raise RuntimeError("loss.backward() must be called before optimizer.step()")
|
||||
model, gradients = optimizer._pending_update
|
||||
optimizer._optimizer.update(model, gradients)
|
||||
mx.eval(model.parameters(), optimizer._optimizer.state, mx.random.state)
|
||||
optimizer._compiled_step(gradients)
|
||||
mx.eval(optimizer._step_state)
|
||||
optimizer._pending_update = None
|
||||
_active_optimizer = None
|
||||
_latest_forward = None
|
||||
|
||||
@@ -24,7 +24,12 @@ def _torch_transpose(self, dim0=None, dim1=None):
|
||||
result = _transpose(self, axes)
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(self, result, lambda value: _transpose(value, axes))
|
||||
return propagate(
|
||||
self,
|
||||
result,
|
||||
lambda value: _transpose(value, axes),
|
||||
("transpose", tuple(axes)),
|
||||
)
|
||||
|
||||
|
||||
def _view(self, *shape):
|
||||
@@ -39,21 +44,36 @@ def _torch_reshape(self, *shape):
|
||||
result = _reshape(self, shape)
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(self, result, lambda value: _reshape(value, shape))
|
||||
return propagate(
|
||||
self,
|
||||
result,
|
||||
lambda value: _reshape(value, shape),
|
||||
("reshape", tuple(shape)),
|
||||
)
|
||||
|
||||
|
||||
def _unsqueeze(self, dim):
|
||||
result = mx.expand_dims(self, axis=dim)
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(self, result, lambda value: mx.expand_dims(value, axis=dim))
|
||||
return propagate(
|
||||
self,
|
||||
result,
|
||||
lambda value: mx.expand_dims(value, axis=dim),
|
||||
("unsqueeze", dim),
|
||||
)
|
||||
|
||||
|
||||
def _torch_squeeze(self, dim=None):
|
||||
result = _squeeze(self, axis=dim)
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(self, result, lambda value: _squeeze(value, axis=dim))
|
||||
return propagate(
|
||||
self,
|
||||
result,
|
||||
lambda value: _squeeze(value, axis=dim),
|
||||
("squeeze", dim),
|
||||
)
|
||||
|
||||
|
||||
def _flatten(self, start_dim=0, end_dim=-1):
|
||||
@@ -188,16 +208,30 @@ def _torch_getitem(self, key):
|
||||
indices = _mask_indices(key)
|
||||
if key.shape == self.shape:
|
||||
result = _getitem(self.reshape(-1), indices)
|
||||
operation = lambda value: _getitem(_reshape(value, (-1,)), indices)
|
||||
operation = lambda value, current_indices: _getitem(
|
||||
_reshape(value, (-1,)), current_indices
|
||||
)
|
||||
signature = ("getitem_bool", "flat")
|
||||
else:
|
||||
result = _getitem(self, indices)
|
||||
operation = lambda value: _getitem(value, indices)
|
||||
operation = lambda value, current_indices: _getitem(
|
||||
value, current_indices
|
||||
)
|
||||
signature = ("getitem_bool", "first_axis")
|
||||
operands = (indices,)
|
||||
elif isinstance(key, mx.array):
|
||||
result = _getitem(self, key)
|
||||
operation = lambda value, current_key: _getitem(value, current_key)
|
||||
signature = ("getitem_array",)
|
||||
operands = (key,)
|
||||
else:
|
||||
result = _getitem(self, key)
|
||||
operation = lambda value: _getitem(value, key)
|
||||
signature = ("getitem", type(key).__qualname__, repr(key))
|
||||
operands = ()
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(self, result, operation)
|
||||
return propagate(self, result, operation, signature, operands)
|
||||
|
||||
|
||||
def _torch_setitem(self, key, value):
|
||||
|
||||
@@ -32,11 +32,11 @@ else:
|
||||
def finish(loss):
|
||||
from torchmlx._autograd import register_loss
|
||||
|
||||
def rebuild(recomputed_input):
|
||||
def rebuild(recomputed_input, current_target, *current_weight):
|
||||
return cross_entropy(
|
||||
recomputed_input,
|
||||
target,
|
||||
weight=weight,
|
||||
current_target,
|
||||
weight=current_weight[0] if current_weight else None,
|
||||
size_average=size_average,
|
||||
ignore_index=ignore_index,
|
||||
reduce=reduce,
|
||||
@@ -44,7 +44,23 @@ else:
|
||||
label_smoothing=label_smoothing,
|
||||
)
|
||||
|
||||
return register_loss(loss, original_input, rebuild)
|
||||
operands = (target,) if weight is None else (target, weight)
|
||||
signature = (
|
||||
"cross_entropy",
|
||||
weight is not None,
|
||||
size_average,
|
||||
ignore_index,
|
||||
reduce,
|
||||
reduction,
|
||||
label_smoothing,
|
||||
)
|
||||
return register_loss(
|
||||
loss,
|
||||
original_input,
|
||||
rebuild,
|
||||
operands,
|
||||
signature,
|
||||
)
|
||||
|
||||
if size_average is not None or reduce is not None:
|
||||
unsupported("torchmlx.nn.functional.cross_entropy legacy reductions")
|
||||
|
||||
@@ -10,6 +10,7 @@ if BACKEND == "torch":
|
||||
return getattr(_native, name)
|
||||
|
||||
else:
|
||||
import mlx.core as mx
|
||||
import mlx.optimizers as _native
|
||||
|
||||
class AdamW:
|
||||
@@ -45,6 +46,19 @@ else:
|
||||
)
|
||||
self._pending_update = None
|
||||
self._random_before_forward = None
|
||||
self._compiled_backward = {}
|
||||
self._optimizer.init(model.trainable_parameters())
|
||||
state = [model.state, self._optimizer.state, mx.random.state]
|
||||
|
||||
def update(gradients):
|
||||
self._optimizer.update(model, gradients)
|
||||
|
||||
self._compiled_step = mx.compile(
|
||||
update,
|
||||
inputs=state,
|
||||
outputs=state,
|
||||
)
|
||||
self._step_state = state
|
||||
|
||||
@property
|
||||
def state(self):
|
||||
|
||||
Reference in new issue
Block a user