compile backward

This commit is contained in:
pj committed 2026-09-20 01:53:35 +05:30
1 parent fefdea2304
commit 69489225c3
6 files changed
+216 -27

No files matched your search

+1
View File
@@ -12,3 +12,4 @@ examples/tinystories-llm/TinyStoriesV2-GPT4-valid.txt
# Environment files # Environment files
.env .env
AGENTS.md
+1 -1
View File
@@ -18,7 +18,7 @@ loss.backward()
optimizer.step() 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. Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation.
+139 -15
View File
@@ -11,6 +11,78 @@ _replayed_losses = {}
_suspended = False _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(): def _snapshot_random_state():
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state] state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
mx.eval(state) mx.eval(state)
@@ -38,7 +110,7 @@ def end_forward(model, args, kwargs, output):
and not _suspended and not _suspended
and isinstance(output, mx.array) and isinstance(output, mx.array)
): ):
_latest_forward = (model, args, kwargs, lambda value: value) _latest_forward = (model, args, kwargs, ())
_lineage[id(output)] = (output, _latest_forward) _lineage[id(output)] = (output, _latest_forward)
@@ -57,37 +129,69 @@ def activate(optimizer):
optimizer._random_before_forward = _snapshot_random_state() 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)) entry = _lineage.get(id(source))
if not _suspended and entry is not None and entry[0] is 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)] = ( _lineage[id(result)] = (
result, result,
( (
model, model,
args, args,
kwargs, kwargs,
lambda output: operation(previous(output)), transforms + ((operation, signature, operands),),
), ),
) )
return result 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: if _suspended or _active_optimizer is None:
return loss return loss
entry = _lineage.get(id(loss_input)) entry = _lineage.get(id(loss_input))
if entry is None or entry[0] is not loss_input: if entry is None or entry[0] is not loss_input:
return loss return loss
model, args, kwargs, transform = entry[1] model, args, kwargs, transforms = entry[1]
def plan(): context = (
def objective(): args,
return rebuild(transform(model(*args, **kwargs))) 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 return loss
@@ -100,12 +204,32 @@ def backward(loss):
raise RuntimeError( raise RuntimeError(
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function" "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() random_after_forward = _snapshot_random_state()
_restore_random_state(_active_optimizer._random_before_forward) _restore_random_state(_active_optimizer._random_before_forward)
_suspended = True _suspended = True
try: 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) mx.eval(replayed_loss, gradients, mx.random.state)
finally: finally:
_suspended = False _suspended = False
@@ -126,8 +250,8 @@ def step(optimizer):
if optimizer._pending_update is None: if optimizer._pending_update is None:
raise RuntimeError("loss.backward() must be called before optimizer.step()") raise RuntimeError("loss.backward() must be called before optimizer.step()")
model, gradients = optimizer._pending_update model, gradients = optimizer._pending_update
optimizer._optimizer.update(model, gradients) optimizer._compiled_step(gradients)
mx.eval(model.parameters(), optimizer._optimizer.state, mx.random.state) mx.eval(optimizer._step_state)
optimizer._pending_update = None optimizer._pending_update = None
_active_optimizer = None _active_optimizer = None
_latest_forward = None _latest_forward = None
+41 -7
View File
@@ -24,7 +24,12 @@ def _torch_transpose(self, dim0=None, dim1=None):
result = _transpose(self, axes) result = _transpose(self, axes)
from ._autograd import propagate 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): def _view(self, *shape):
@@ -39,21 +44,36 @@ def _torch_reshape(self, *shape):
result = _reshape(self, shape) result = _reshape(self, shape)
from ._autograd import propagate 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): def _unsqueeze(self, dim):
result = mx.expand_dims(self, axis=dim) result = mx.expand_dims(self, axis=dim)
from ._autograd import propagate 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): def _torch_squeeze(self, dim=None):
result = _squeeze(self, axis=dim) result = _squeeze(self, axis=dim)
from ._autograd import propagate 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): def _flatten(self, start_dim=0, end_dim=-1):
@@ -188,16 +208,30 @@ def _torch_getitem(self, key):
indices = _mask_indices(key) indices = _mask_indices(key)
if key.shape == self.shape: if key.shape == self.shape:
result = _getitem(self.reshape(-1), indices) 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: else:
result = _getitem(self, indices) 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: else:
result = _getitem(self, key) result = _getitem(self, key)
operation = lambda value: _getitem(value, key) operation = lambda value: _getitem(value, key)
signature = ("getitem", type(key).__qualname__, repr(key))
operands = ()
from ._autograd import propagate from ._autograd import propagate
return propagate(self, result, operation) return propagate(self, result, operation, signature, operands)
def _torch_setitem(self, key, value): def _torch_setitem(self, key, value):
+20 -4
View File
@@ -32,11 +32,11 @@ else:
def finish(loss): def finish(loss):
from torchmlx._autograd import register_loss from torchmlx._autograd import register_loss
def rebuild(recomputed_input): def rebuild(recomputed_input, current_target, *current_weight):
return cross_entropy( return cross_entropy(
recomputed_input, recomputed_input,
target, current_target,
weight=weight, weight=current_weight[0] if current_weight else None,
size_average=size_average, size_average=size_average,
ignore_index=ignore_index, ignore_index=ignore_index,
reduce=reduce, reduce=reduce,
@@ -44,7 +44,23 @@ else:
label_smoothing=label_smoothing, 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: if size_average is not None or reduce is not None:
unsupported("torchmlx.nn.functional.cross_entropy legacy reductions") unsupported("torchmlx.nn.functional.cross_entropy legacy reductions")
+14
View File
@@ -10,6 +10,7 @@ if BACKEND == "torch":
return getattr(_native, name) return getattr(_native, name)
else: else:
import mlx.core as mx
import mlx.optimizers as _native import mlx.optimizers as _native
class AdamW: class AdamW:
@@ -45,6 +46,19 @@ else:
) )
self._pending_update = None self._pending_update = None
self._random_before_forward = 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 @property
def state(self): def state(self):