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() 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. 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 _active_optimizer = None
_forward_depth = 0 _forward_depth = 0
_latest_forward = None
_loss_plans = {} _loss_plans = {}
_lineage = {} _lineage = {}
_replayed_losses = {} _replayed_losses = {}
@@ -16,24 +15,48 @@ class _ArraySlot:
self.index = index self.index = index
class _BackwardPlan: class _TrainingPlan:
def __init__(self, value_and_grad, state): def __init__(self, model, optimizer, objective, compile=True):
self._value_and_grad = value_and_grad self.model = model
self._compiled = mx.compile( self.optimizer = optimizer
value_and_grad, self._value_and_grad = nn.value_and_grad(model, objective)
inputs=state, self._compile_enabled = compile
outputs=state, 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): def _fused(self, *inputs):
function = self._compiled if self._compile_enabled else self._value_and_grad loss, gradients = self._value_and_grad(*inputs)
return function(*inputs) self.optimizer.update(self.model, gradients)
return loss
def fallback(self, *inputs): def backward(self, *inputs):
self._compile_enabled = False if self._compile_enabled:
return self._compiled_backward(*inputs)
return self._value_and_grad(*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): def _partition_arrays(value, arrays):
if isinstance(value, mx.array): if isinstance(value, mx.array):
@@ -45,9 +68,7 @@ def _partition_arrays(value, arrays):
if isinstance(value, list): if isinstance(value, list):
return [_partition_arrays(item, arrays) for item in value] return [_partition_arrays(item, arrays) for item in value]
if isinstance(value, dict): if isinstance(value, dict):
return { return {key: _partition_arrays(item, arrays) for key, item in value.items()}
key: _partition_arrays(item, arrays) for key, item in value.items()
}
return value return value
@@ -59,9 +80,7 @@ def _restore_arrays(template, arrays):
if isinstance(template, list): if isinstance(template, list):
return [_restore_arrays(item, arrays) for item in template] return [_restore_arrays(item, arrays) for item in template]
if isinstance(template, dict): if isinstance(template, dict):
return { return {key: _restore_arrays(item, arrays) for key, item in template.items()}
key: _restore_arrays(item, arrays) for key, item in template.items()
}
return template return template
@@ -83,6 +102,16 @@ def _template_key(template):
return ("static", type(template).__qualname__, repr(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(): 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)
@@ -95,23 +124,54 @@ def _restore_random_state(state):
mx.eval(mx.random.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 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 _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): def end_forward(model, args, kwargs, output):
global _forward_depth, _latest_forward global _forward_depth
_forward_depth -= 1 _forward_depth -= 1
if ( if _forward_depth != 0 or not isinstance(output, mx.array):
_forward_depth == 0 return
and _active_optimizer is not None random_before = _active_optimizer._random_before_forward
and model is _active_optimizer._model stochastic = any(
and not _suspended current is not previous
and isinstance(output, mx.array) for current, previous in zip(mx.random.state, random_before)
): )
_latest_forward = (model, args, kwargs, ()) descriptor = (
_lineage[id(output)] = (output, _latest_forward) 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(): def abort_forward():
@@ -120,19 +180,24 @@ def abort_forward():
def activate(optimizer): def activate(optimizer):
global _active_optimizer, _latest_forward global _active_optimizer, _forward_depth
_active_optimizer = optimizer _active_optimizer = optimizer
_latest_forward = None _forward_depth = 0
_loss_plans.clear() _loss_plans.clear()
_lineage.clear() _lineage.clear()
_replayed_losses.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=()): def propagate(source, result, operation, signature, operands=()):
if _suspended:
return result
entry = _lineage.get(id(source)) entry = _lineage.get(id(source))
if not _suspended and entry is not None and entry[0] is source: if entry is None or entry[0] is not source:
model, args, kwargs, transforms = entry[1] return result
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
_lineage[id(result)] = ( _lineage[id(result)] = (
result, result,
( (
@@ -140,19 +205,45 @@ def propagate(source, result, operation, signature, operands=()):
args, args,
kwargs, kwargs,
transforms + ((operation, signature, operands),), transforms + ((operation, signature, operands),),
random_before,
compile_allowed,
), ),
) )
return result 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): 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, transforms = entry[1] model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
context = ( context = (
args, args,
kwargs, kwargs,
@@ -166,37 +257,23 @@ def register_loss(loss, loss_input, rebuild, operands, signature):
_template_key(template), _template_key(template),
tuple(transform_signature for _, transform_signature, _ in transforms), tuple(transform_signature for _, transform_signature, _ in transforms),
signature, 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_plans[id(loss)] = (
loss, loss,
model, model,
cache_key, cache_key,
tuple(dynamic_inputs), tuple(dynamic_inputs),
build, template,
transforms,
rebuild,
random_before,
compile_allowed,
) )
return loss return loss
def backward(loss): def backward(loss):
global _suspended
if _active_optimizer is None: if _active_optimizer is None:
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()") raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
entry = _loss_plans.get(id(loss)) entry = _loss_plans.get(id(loss))
@@ -204,56 +281,229 @@ 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, 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 _suspended = True
try: 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: try:
replayed_loss, gradients = plan(*dynamic_inputs) result = plan.fused(*inputs) if fused else plan.backward(*inputs)
except ValueError as error: current_parameters = model.parameters()
message = str(error) current_optimizer_state = _copy_tree(optimizer.state) if fused else []
if ( mx.eval(
not plan._compile_enabled result,
or "Attempting to eval an array" not in message current_parameters,
or not ( current_optimizer_state,
"function transformations" in message mx.random.state
or "without a primitive" in message 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 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: finally:
_suspended = False _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) _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): 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)) entry = _replayed_losses.get(id(value))
if entry is None or entry[0] is not value: if entry is None or entry[0] is not value:
return value return value
return entry[1] return entry[1]
def step(optimizer): def _finish_step(loss, replayed_loss):
global _active_optimizer, _latest_forward global _active_optimizer, _forward_depth
if optimizer._pending_update is None: _replayed_losses[id(loss)] = (loss, replayed_loss)
raise RuntimeError("loss.backward() must be called before optimizer.step()") _active_optimizer._pending_update = None
model, gradients = optimizer._pending_update
optimizer._compiled_step(gradients)
mx.eval(optimizer._step_state)
optimizer._pending_update = None
_active_optimizer = None _active_optimizer = None
_latest_forward = None _forward_depth = 0
_loss_plans.clear() _loss_plans.clear()
_lineage.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): 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)) return _item(replayed_value(self))
+3 -1
View File
@@ -32,12 +32,14 @@ else:
def __call__(self, *args, **kwargs): def __call__(self, *args, **kwargs):
from torchmlx._autograd import abort_forward, begin_forward, end_forward from torchmlx._autograd import abort_forward, begin_forward, end_forward
begin_forward() tracking = begin_forward(self)
try: try:
output = self.forward(*args, **kwargs) output = self.forward(*args, **kwargs)
except Exception: except Exception:
if tracking:
abort_forward() abort_forward()
raise raise
if tracking:
end_forward(self, args, kwargs, output) end_forward(self, args, kwargs, output)
return output return output
+1 -1
View File
@@ -46,7 +46,7 @@ else:
) )
self._pending_update = None self._pending_update = None
self._random_before_forward = None self._random_before_forward = None
self._compiled_backward = {} self._training_plans = {}
self._optimizer.init(model.trainable_parameters()) self._optimizer.init(model.trainable_parameters())
state = [model.state, self._optimizer.state, mx.random.state] 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) return self._step(x, y)
else: else:
import mlx.core as mx
import mlx.nn as nn
class Trainer: class Trainer:
def __init__(self, model, optimizer, loss_fn, compile=True): def __init__(self, model, optimizer, loss_fn, compile=True):
self.model = model self.model = model
self.optimizer = optimizer self.optimizer = optimizer
self.loss_fn = loss_fn self.loss_fn = loss_fn
native_optimizer = optimizer._optimizer self.compile = compile
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
def step(self, x, y): def step(self, x, y):
loss = self._step(x, y) from ._autograd import trainer_step
mx.eval(
loss, return trainer_step(self, x, y)
self.model.parameters(),
self.optimizer._optimizer.state,
mx.random.state,
)
return loss