fuse mlx training step

This commit is contained in:
pj committed 2026-09-20 14:09:39 +05:30
1 parent 69489225c3
commit 4eb86111a8
6 files changed
+352 -124

No files matched your search

+1 -1
View File
@@ -18,7 +18,7 @@ loss.backward()
optimizer.step()
```
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.
MLX implements this sequence by recording the outer model call and replaying the forward, backward, and AdamW update inside one cached compiled graph during `step`. 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. Reading `loss.item()` before `step` materializes gradients through a separate compiled path. 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.
+341 -91
View File
@@ -4,7 +4,6 @@ import mlx.nn as nn
_active_optimizer = None
_forward_depth = 0
_latest_forward = None
_loss_plans = {}
_lineage = {}
_replayed_losses = {}
@@ -16,24 +15,48 @@ class _ArraySlot:
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,
class _TrainingPlan:
def __init__(self, model, optimizer, objective, compile=True):
self.model = model
self.optimizer = optimizer
self._value_and_grad = nn.value_and_grad(model, objective)
self._compile_enabled = compile
self._compiled_backward = None
self._compiled_fused = None
self._backward_signatures = set()
self._fused_signatures = set()
if compile:
backward_state = [model.state, mx.random.state]
fused_state = [model.state, optimizer.state, mx.random.state]
self._compiled_backward = mx.compile(
self._value_and_grad,
inputs=backward_state,
outputs=backward_state,
)
self._compiled_fused = mx.compile(
self._fused,
inputs=fused_state,
outputs=fused_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 _fused(self, *inputs):
loss, gradients = self._value_and_grad(*inputs)
self.optimizer.update(self.model, gradients)
return loss
def fallback(self, *inputs):
self._compile_enabled = False
def backward(self, *inputs):
if self._compile_enabled:
return self._compiled_backward(*inputs)
return self._value_and_grad(*inputs)
def fused(self, *inputs):
if self._compile_enabled:
return self._compiled_fused(*inputs)
return self._fused(*inputs)
def disable_compile(self):
self._compile_enabled = False
def _partition_arrays(value, arrays):
if isinstance(value, mx.array):
@@ -45,9 +68,7 @@ def _partition_arrays(value, arrays):
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 {key: _partition_arrays(item, arrays) for key, item in value.items()}
return value
@@ -59,9 +80,7 @@ def _restore_arrays(template, arrays):
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 {key: _restore_arrays(item, arrays) for key, item in template.items()}
return template
@@ -83,6 +102,16 @@ def _template_key(template):
return ("static", type(template).__qualname__, repr(template))
def _copy_tree(value):
if isinstance(value, dict):
return {key: _copy_tree(item) for key, item in value.items()}
if isinstance(value, list):
return [_copy_tree(item) for item in value]
if isinstance(value, tuple):
return tuple(_copy_tree(item) for item in value)
return value
def _snapshot_random_state():
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
mx.eval(state)
@@ -95,23 +124,54 @@ def _restore_random_state(state):
mx.eval(mx.random.state)
def begin_forward():
def _compile_failure(error):
message = str(error)
return "Attempting to eval an array" in message and (
"function transformations" in message or "without a primitive" in message
)
def begin_forward(model):
global _forward_depth
if (
_suspended
or _active_optimizer is None
or model is not _active_optimizer._model
):
return False
if _forward_depth:
_active_optimizer._recursive_forward = True
_forward_depth += 1
return True
def note_eager_evaluation():
if _forward_depth and _active_optimizer is not None and not _suspended:
_active_optimizer._eager_forward = True
def end_forward(model, args, kwargs, output):
global _forward_depth, _latest_forward
global _forward_depth
_forward_depth -= 1
if (
_forward_depth == 0
and _active_optimizer is not None
and model is _active_optimizer._model
and not _suspended
and isinstance(output, mx.array)
):
_latest_forward = (model, args, kwargs, ())
_lineage[id(output)] = (output, _latest_forward)
if _forward_depth != 0 or not isinstance(output, mx.array):
return
random_before = _active_optimizer._random_before_forward
stochastic = any(
current is not previous
for current, previous in zip(mx.random.state, random_before)
)
descriptor = (
model,
args,
kwargs,
(),
random_before if stochastic else None,
not (
_active_optimizer._recursive_forward
or _active_optimizer._eager_forward
),
)
_lineage[id(output)] = (output, descriptor)
def abort_forward():
@@ -120,19 +180,24 @@ def abort_forward():
def activate(optimizer):
global _active_optimizer, _latest_forward
global _active_optimizer, _forward_depth
_active_optimizer = optimizer
_latest_forward = None
_forward_depth = 0
_loss_plans.clear()
_lineage.clear()
_replayed_losses.clear()
optimizer._random_before_forward = _snapshot_random_state()
optimizer._random_before_forward = tuple(mx.random.state)
optimizer._recursive_forward = False
optimizer._eager_forward = False
def propagate(source, result, operation, signature, operands=()):
if _suspended:
return result
entry = _lineage.get(id(source))
if not _suspended and entry is not None and entry[0] is source:
model, args, kwargs, transforms = entry[1]
if entry is None or entry[0] is not source:
return result
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
_lineage[id(result)] = (
result,
(
@@ -140,19 +205,45 @@ def propagate(source, result, operation, signature, operands=()):
args,
kwargs,
transforms + ((operation, signature, operands),),
random_before,
compile_allowed,
),
)
return result
def _plan_for(
optimizer, model, cache_key, template, transforms, rebuild, compile_allowed
):
plan = optimizer._training_plans.get(cache_key)
if plan is not None:
return 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)
plan = _TrainingPlan(
model, optimizer._optimizer, objective, compile=compile_allowed
)
optimizer._training_plans[cache_key] = plan
return plan
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, transforms = entry[1]
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
context = (
args,
kwargs,
@@ -166,37 +257,23 @@ def register_loss(loss, loss_input, rebuild, operands, signature):
_template_key(template),
tuple(transform_signature for _, transform_signature, _ in transforms),
signature,
compile_allowed,
)
def build():
state = [model.state, mx.random.state]
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,
template,
transforms,
rebuild,
random_before,
compile_allowed,
)
return loss
def backward(loss):
global _suspended
if _active_optimizer is None:
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
entry = _loss_plans.get(id(loss))
@@ -204,56 +281,229 @@ def backward(loss):
raise RuntimeError(
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
)
_, model, cache_key, dynamic_inputs, build = entry
random_after_forward = _snapshot_random_state()
_restore_random_state(_active_optimizer._random_before_forward)
(
_,
model,
cache_key,
dynamic_inputs,
template,
transforms,
rebuild,
random_before,
compile_allowed,
) = entry
plan = _plan_for(
_active_optimizer,
model,
cache_key,
template,
transforms,
rebuild,
compile_allowed,
)
_active_optimizer._pending_update = (
"staged",
loss,
plan,
dynamic_inputs,
random_before,
)
def _restore_execution(model, optimizer, parameters, optimizer_state):
model.update(parameters)
if optimizer_state is not None:
_restore_tree(optimizer.state, optimizer_state)
def _restore_tree(current, saved):
if isinstance(current, dict) and isinstance(saved, dict):
for key in tuple(current):
if key not in saved:
del current[key]
for key, value in saved.items():
if key in current:
restored = _restore_tree(current[key], value)
if restored is not current[key]:
current[key] = restored
else:
current[key] = value
return current
if isinstance(current, list) and isinstance(saved, list):
restored = [
_restore_tree(old, value) for old, value in zip(current, saved)
]
current[:] = restored + saved[len(restored) :]
return current
return saved
def _execute(plan, inputs, fused, random_before, advance_random=False):
global _suspended
model = plan.model
optimizer = plan.optimizer
signatures = plan._fused_signatures if fused else plan._backward_signatures
signature = tuple((value.shape, value.dtype) for value in inputs)
protect_state = not plan._compile_enabled or signature not in signatures
parameters = (
_copy_tree(model.trainable_parameters())
if protect_state
else getattr(optimizer, "_torchmlx_parameters", None)
)
optimizer_state = None
if fused:
optimizer_state = (
_copy_tree(optimizer.state)
if protect_state
else getattr(optimizer, "_torchmlx_state", None)
)
random_after = _snapshot_random_state() if random_before is not None else None
random_rollback = tuple(mx.random.state) if advance_random else None
if random_before is not None:
_restore_random_state(random_before)
_suspended = True
try:
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
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
current_parameters = model.parameters()
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
mx.eval(
result,
current_parameters,
current_optimizer_state,
mx.random.state
if random_before is not None or advance_random
else [],
)
):
signatures.add(signature)
if fused:
optimizer._torchmlx_parameters = current_parameters
optimizer._torchmlx_state = current_optimizer_state
except (ValueError, RuntimeError) as error:
if not plan._compile_enabled or not _compile_failure(error):
raise
_restore_execution(model, optimizer, parameters, optimizer_state)
if random_before is not None:
_restore_random_state(random_before)
elif random_rollback is not None:
_restore_random_state(random_rollback)
plan.disable_compile()
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
current_parameters = model.parameters()
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
mx.eval(
result,
current_parameters,
current_optimizer_state,
mx.random.state
if random_before is not None or advance_random
else [],
)
if fused:
optimizer._torchmlx_parameters = current_parameters
optimizer._torchmlx_state = current_optimizer_state
except Exception:
if parameters is not None:
_restore_execution(model, optimizer, parameters, optimizer_state)
if random_rollback is not None:
_restore_random_state(random_rollback)
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
_restore_random_state(random_after_forward)
if random_after is not None:
_restore_random_state(random_after)
return result
def _materialize(optimizer):
pending = optimizer._pending_update
_, loss, plan, dynamic_inputs, random_before = pending
replayed_loss, gradients = _execute(
plan, dynamic_inputs, fused=False, random_before=random_before
)
_replayed_losses[id(loss)] = (loss, replayed_loss)
_active_optimizer._pending_update = (model, gradients)
optimizer._pending_update = ("materialized", loss, plan.model, gradients)
return replayed_loss
def replayed_value(value):
if _active_optimizer is not None:
pending = _active_optimizer._pending_update
if pending is not None and pending[0] == "staged" and pending[1] is value:
return _materialize(_active_optimizer)
entry = _replayed_losses.get(id(value))
if entry is None or entry[0] is not value:
return value
return entry[1]
def step(optimizer):
global _active_optimizer, _latest_forward
if optimizer._pending_update is None:
raise RuntimeError("loss.backward() must be called before optimizer.step()")
model, gradients = optimizer._pending_update
optimizer._compiled_step(gradients)
mx.eval(optimizer._step_state)
optimizer._pending_update = None
def _finish_step(loss, replayed_loss):
global _active_optimizer, _forward_depth
_replayed_losses[id(loss)] = (loss, replayed_loss)
_active_optimizer._pending_update = None
_active_optimizer = None
_latest_forward = None
_forward_depth = 0
_loss_plans.clear()
_lineage.clear()
def step(optimizer):
pending = optimizer._pending_update
if pending is None:
raise RuntimeError("loss.backward() must be called before optimizer.step()")
if pending[0] == "staged":
_, loss, plan, dynamic_inputs, random_before = pending
replayed_loss = _execute(
plan, dynamic_inputs, fused=True, random_before=random_before
)
else:
_, loss, model, gradients = pending
parameters = _copy_tree(model.trainable_parameters())
optimizer_state = _copy_tree(optimizer._optimizer.state)
try:
optimizer._compiled_step(gradients)
current_parameters = model.parameters()
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
mx.eval(current_parameters, current_optimizer_state)
except Exception:
_restore_execution(
model, optimizer._optimizer, parameters, optimizer_state
)
raise
optimizer._optimizer._torchmlx_parameters = current_parameters
optimizer._optimizer._torchmlx_state = current_optimizer_state
replayed_loss = _replayed_losses[id(loss)][1]
_finish_step(loss, replayed_loss)
def trainer_step(trainer, x, y):
context = (x, y)
dynamic_inputs = []
template = _partition_arrays(context, dynamic_inputs)
cache_key = (
getattr(trainer.model, "training", None),
_template_key(template),
(),
("trainer", id(trainer.loss_fn), trainer.compile),
)
plan = trainer.optimizer._training_plans.get(cache_key)
if plan is None:
def objective(*current_inputs):
current_x, current_y = _restore_arrays(template, current_inputs)
return trainer.loss_fn(trainer.model(current_x), current_y)
plan = _TrainingPlan(
trainer.model,
trainer.optimizer._optimizer,
objective,
compile=trainer.compile,
)
trainer.optimizer._training_plans[cache_key] = plan
return _execute(
plan,
tuple(dynamic_inputs),
fused=True,
random_before=None,
advance_random=True,
)
+2 -1
View File
@@ -189,8 +189,9 @@ def _backward(self, *args, **kwargs):
def _torch_item(self):
from ._autograd import replayed_value
from ._autograd import note_eager_evaluation, replayed_value
note_eager_evaluation()
return _item(replayed_value(self))
+3 -1
View File
@@ -32,12 +32,14 @@ else:
def __call__(self, *args, **kwargs):
from torchmlx._autograd import abort_forward, begin_forward, end_forward
begin_forward()
tracking = begin_forward(self)
try:
output = self.forward(*args, **kwargs)
except Exception:
if tracking:
abort_forward()
raise
if tracking:
end_forward(self, args, kwargs, output)
return output
+1 -1
View File
@@ -46,7 +46,7 @@ else:
)
self._pending_update = None
self._random_before_forward = None
self._compiled_backward = {}
self._training_plans = {}
self._optimizer.init(model.trainable_parameters())
state = [model.state, self._optimizer.state, mx.random.state]
+4 -29
View File
@@ -22,39 +22,14 @@ if BACKEND == "torch":
return self._step(x, y)
else:
import mlx.core as mx
import mlx.nn as nn
class Trainer:
def __init__(self, model, optimizer, loss_fn, compile=True):
self.model = model
self.optimizer = optimizer
self.loss_fn = loss_fn
native_optimizer = optimizer._optimizer
native_optimizer.init(model.trainable_parameters())
def objective(x, y):
return loss_fn(model(x), y)
value_and_grad = nn.value_and_grad(model, objective)
def train_step(x, y):
loss, gradients = value_and_grad(x, y)
native_optimizer.update(model, gradients)
return loss
if compile:
state = [model.state, native_optimizer.state, mx.random.state]
self._step = mx.compile(train_step, inputs=state, outputs=state)
else:
self._step = train_step
self.compile = compile
def step(self, x, y):
loss = self._step(x, y)
mx.eval(
loss,
self.model.parameters(),
self.optimizer._optimizer.state,
mx.random.state,
)
return loss
from ._autograd import trainer_step
return trainer_step(self, x, y)