support backward training

This commit is contained in:
pj committed 2026-09-19 23:10:37 +05:30
1 parent 7fa8ae82ad
commit 53c9cf4159
7 files changed
+277 -18

No files matched your search

+13 -1
View File
@@ -2,12 +2,24 @@
TorchMLX targets the common transformer operations used by GPT-2, Llama 3, Qwen 3, and GPT-OSS style implementations. TorchMLX targets the common transformer operations used by GPT-2, Llama 3, Qwen 3, and GPT-OSS style implementations.
Supported MLX operations include embeddings, linear layers, normalization building blocks, dropout, activations, causal attention, tensor shape operations, masks, top-k routing, and AdamW training through `Trainer`. Supported MLX operations include embeddings, linear layers, normalization building blocks, dropout, activations, causal attention, tensor shape operations, masks, top-k routing, and AdamW training.
MLX arrays remain native arrays. Torch-style tensor methods are installed on the native array type for the supported subset. MLX arrays remain native arrays. Torch-style tensor methods are installed on the native array type for the supported subset.
Boolean expert routing and `unique` execute eagerly because their output shapes control Python flow. Boolean expert routing and `unique` execute eagerly because their output shapes control Python flow.
The standard training sequence works with `cross_entropy` losses applied directly or after reshape, view, slicing, transpose, squeeze, unsqueeze, or flatten operations:
```python
optimizer.zero_grad()
logits = model(input_tokens)
loss = F.cross_entropy(logits.reshape(-1, vocabulary_size), targets.reshape(-1))
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.
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.
The referenced OpenArch model files contain source errors independent of TorchMLX, including invalid constructor calls and undefined attributes. Correct those errors before using either backend. The referenced OpenArch model files contain source errors independent of TorchMLX, including invalid constructor calls and undefined attributes. Correct those errors before using either backend.
+6 -2
View File
@@ -235,13 +235,17 @@ config = TransformerConfig(vocabulary_size=len(tokenizer))
encoded_text = tokenizer.encode(training_text) encoded_text = tokenizer.encode(training_text)
model = TinyStoriesTransformer(config).to(device) model = TinyStoriesTransformer(config).to(device)
optimizer = optim.AdamW(model.parameters(), lr=3e-4) optimizer = optim.AdamW(model.parameters(), lr=3e-4)
trainer = torch.Trainer(model, optimizer, language_model_loss, compile=True) model.train()
for step in range(arguments.steps): for step in range(arguments.steps):
input_tokens, target_tokens = create_batch( input_tokens, target_tokens = create_batch(
encoded_text, batch_size=8, context_length=config.context_length encoded_text, batch_size=8, context_length=config.context_length
) )
loss = trainer.step(input_tokens, target_tokens) optimizer.zero_grad()
logits = model(input_tokens)
loss = language_model_loss(logits, target_tokens)
loss.backward()
optimizer.step()
print(f"step {step + 1}: loss {loss.item():.4f}") print(f"step {step + 1}: loss {loss.item():.4f}")
print(generate_text(model, tokenizer, "Once upon a time", token_count=120)) print(generate_text(model, tokenizer, "Once upon a time", token_count=120))
+135
View File
@@ -0,0 +1,135 @@
import mlx.core as mx
import mlx.nn as nn
_active_optimizer = None
_forward_depth = 0
_latest_forward = None
_loss_plans = {}
_lineage = {}
_replayed_losses = {}
_suspended = False
def _snapshot_random_state():
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
mx.eval(state)
return state
def _restore_random_state(state):
for current_key, saved_key in zip(mx.random.state, state):
current_key[...] = saved_key
mx.eval(mx.random.state)
def begin_forward():
global _forward_depth
_forward_depth += 1
def end_forward(model, args, kwargs, output):
global _forward_depth, _latest_forward
_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, lambda value: value)
_lineage[id(output)] = (output, _latest_forward)
def abort_forward():
global _forward_depth
_forward_depth -= 1
def activate(optimizer):
global _active_optimizer, _latest_forward
_active_optimizer = optimizer
_latest_forward = None
_loss_plans.clear()
_lineage.clear()
_replayed_losses.clear()
optimizer._random_before_forward = _snapshot_random_state()
def propagate(source, result, operation=lambda value: value):
entry = _lineage.get(id(source))
if not _suspended and entry is not None and entry[0] is source:
model, args, kwargs, previous = entry[1]
_lineage[id(result)] = (
result,
(
model,
args,
kwargs,
lambda output: operation(previous(output)),
),
)
return result
def register_loss(loss, loss_input, rebuild):
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]
def plan():
def objective():
return rebuild(transform(model(*args, **kwargs)))
return nn.value_and_grad(model, objective)()
_loss_plans[id(loss)] = (loss, model, plan)
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))
if entry is None or entry[0] is not loss:
raise RuntimeError(
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
)
_, model, plan = entry
random_after_forward = _snapshot_random_state()
_restore_random_state(_active_optimizer._random_before_forward)
_suspended = True
try:
replayed_loss, gradients = plan()
mx.eval(replayed_loss, gradients, mx.random.state)
finally:
_suspended = False
_restore_random_state(random_after_forward)
_replayed_losses[id(loss)] = (loss, replayed_loss)
_active_optimizer._pending_update = (model, gradients)
def replayed_value(value):
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._optimizer.update(model, gradients)
mx.eval(model.parameters(), optimizer._optimizer.state, mx.random.state)
optimizer._pending_update = None
_active_optimizer = None
_latest_forward = None
_loss_plans.clear()
_lineage.clear()
+64 -9
View File
@@ -2,11 +2,14 @@ import mlx.core as mx
_transpose = mx.array.transpose _transpose = mx.array.transpose
_reshape = mx.array.reshape
_squeeze = mx.array.squeeze
_mean = mx.array.mean _mean = mx.array.mean
_var = mx.array.var _var = mx.array.var
_any = mx.array.any _any = mx.array.any
_getitem = mx.array.__getitem__ _getitem = mx.array.__getitem__
_setitem = mx.array.__setitem__ _setitem = mx.array.__setitem__
_item = mx.array.item
def _torch_transpose(self, dim0=None, dim1=None): def _torch_transpose(self, dim0=None, dim1=None):
@@ -18,17 +21,49 @@ def _torch_transpose(self, dim0=None, dim1=None):
else: else:
axes = list(range(self.ndim)) axes = list(range(self.ndim))
axes[dim0], axes[dim1] = axes[dim1], axes[dim0] axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
return _transpose(self, axes) result = _transpose(self, axes)
from ._autograd import propagate
return propagate(self, result, lambda value: _transpose(value, axes))
def _view(self, *shape): def _view(self, *shape):
if len(shape) == 1 and isinstance(shape[0], (tuple, list)): if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
shape = shape[0] shape = shape[0]
return self.reshape(shape) return _torch_reshape(self, shape)
def _torch_reshape(self, *shape):
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
shape = shape[0]
result = _reshape(self, shape)
from ._autograd import propagate
return propagate(self, result, lambda value: _reshape(value, shape))
def _unsqueeze(self, dim): def _unsqueeze(self, dim):
return mx.expand_dims(self, axis=dim) result = mx.expand_dims(self, axis=dim)
from ._autograd import propagate
return propagate(self, result, lambda value: mx.expand_dims(value, axis=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))
def _flatten(self, start_dim=0, end_dim=-1):
if end_dim < 0:
end_dim += self.ndim
flattened = 1
for dimension in self.shape[start_dim : end_dim + 1]:
flattened *= dimension
shape = self.shape[:start_dim] + (flattened,) + self.shape[end_dim + 1 :]
return _torch_reshape(self, shape)
def _float(self): def _float(self):
@@ -126,9 +161,17 @@ def _requires_grad(self, requires_grad=True):
def _backward(self, *args, **kwargs): def _backward(self, *args, **kwargs):
raise RuntimeError( if args or kwargs:
"loss.backward() is not supported by the MLX backend; use torchmlx.Trainer" raise TypeError("MLX backward compatibility does not accept arguments")
) from ._autograd import backward
backward(self)
def _torch_item(self):
from ._autograd import replayed_value
return _item(replayed_value(self))
def _mask_indices(mask): def _mask_indices(mask):
@@ -144,9 +187,17 @@ def _torch_getitem(self, key):
if isinstance(key, mx.array) and key.dtype == mx.bool_: if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key) indices = _mask_indices(key)
if key.shape == self.shape: if key.shape == self.shape:
return _getitem(self.reshape(-1), indices) result = _getitem(self.reshape(-1), indices)
return _getitem(self, indices) operation = lambda value: _getitem(_reshape(value, (-1,)), indices)
return _getitem(self, key) else:
result = _getitem(self, indices)
operation = lambda value: _getitem(value, indices)
else:
result = _getitem(self, key)
operation = lambda value: _getitem(value, key)
from ._autograd import propagate
return propagate(self, result, operation)
def _torch_setitem(self, key, value): def _torch_setitem(self, key, value):
@@ -163,8 +214,11 @@ def _torch_setitem(self, key, value):
def install(device_type): def install(device_type):
mx.array.transpose = _torch_transpose mx.array.transpose = _torch_transpose
mx.array.reshape = _torch_reshape
mx.array.view = _view mx.array.view = _view
mx.array.unsqueeze = _unsqueeze mx.array.unsqueeze = _unsqueeze
mx.array.squeeze = _torch_squeeze
mx.array.flatten = _flatten
mx.array.float = _float mx.array.float = _float
mx.array.bool = _bool mx.array.bool = _bool
mx.array.pow = _pow mx.array.pow = _pow
@@ -180,6 +234,7 @@ def install(device_type):
mx.array.contiguous = _contiguous mx.array.contiguous = _contiguous
mx.array.requires_grad_ = _requires_grad mx.array.requires_grad_ = _requires_grad
mx.array.backward = _backward mx.array.backward = _backward
mx.array.item = _torch_item
mx.array.device = property(lambda self: device_type("mps")) mx.array.device = property(lambda self: device_type("mps"))
mx.array.__getitem__ = _torch_getitem mx.array.__getitem__ = _torch_getitem
mx.array.__setitem__ = _torch_setitem mx.array.__setitem__ = _torch_setitem
+18 -1
View File
@@ -23,15 +23,32 @@ else:
import mlx.core as mx import mlx.core as mx
import mlx.nn as _native import mlx.nn as _native
class _ParameterTree(dict):
def __init__(self, values, model):
super().__init__(values)
self.model = model
class Module(_native.Module): class Module(_native.Module):
def __call__(self, *args, **kwargs): def __call__(self, *args, **kwargs):
return self.forward(*args, **kwargs) from torchmlx._autograd import abort_forward, begin_forward, end_forward
begin_forward()
try:
output = self.forward(*args, **kwargs)
except Exception:
abort_forward()
raise
end_forward(self, args, kwargs, output)
return output
def forward(self, *args, **kwargs): def forward(self, *args, **kwargs):
raise NotImplementedError( raise NotImplementedError(
f"Module [{type(self).__name__}] is missing the required forward function" f"Module [{type(self).__name__}] is missing the required forward function"
) )
def parameters(self):
return _ParameterTree(super().parameters(), self)
def to(self, *args, **kwargs): def to(self, *args, **kwargs):
dtype = kwargs.pop("dtype", None) dtype = kwargs.pop("dtype", None)
device = kwargs.pop("device", None) device = kwargs.pop("device", None)
+22 -3
View File
@@ -27,6 +27,25 @@ else:
reduction="mean", reduction="mean",
label_smoothing=0.0, label_smoothing=0.0,
): ):
original_input = input
def finish(loss):
from torchmlx._autograd import register_loss
def rebuild(recomputed_input):
return cross_entropy(
recomputed_input,
target,
weight=weight,
size_average=size_average,
ignore_index=ignore_index,
reduce=reduce,
reduction=reduction,
label_smoothing=label_smoothing,
)
return register_loss(loss, original_input, rebuild)
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")
if input.ndim < 2: if input.ndim < 2:
@@ -56,15 +75,15 @@ else:
losses = (1 - label_smoothing) * target_losses + label_smoothing * smooth_losses losses = (1 - label_smoothing) * target_losses + label_smoothing * smooth_losses
losses = mx.where(valid, losses, mx.zeros_like(losses)) losses = mx.where(valid, losses, mx.zeros_like(losses))
if reduction == "none": if reduction == "none":
return losses.reshape(output_shape) return finish(losses.reshape(output_shape))
if reduction == "sum": if reduction == "sum":
return mx.sum(losses) return finish(mx.sum(losses))
if reduction == "mean": if reduction == "mean":
if weight is None: if weight is None:
denominator = mx.sum(valid) denominator = mx.sum(valid)
else: else:
denominator = mx.sum(mx.where(valid, weight[safe_targets], 0)) denominator = mx.sum(mx.where(valid, weight[safe_targets], 0))
return mx.sum(losses) / denominator return finish(mx.sum(losses) / denominator)
raise ValueError(f"invalid reduction {reduction!r}") raise ValueError(f"invalid reduction {reduction!r}")
def scaled_dot_product_attention( def scaled_dot_product_attention(
+19 -2
View File
@@ -29,7 +29,13 @@ else:
): ):
if amsgrad or maximize or foreach is not None or capturable or differentiable or fused is not None: if amsgrad or maximize or foreach is not None or capturable or differentiable or fused is not None:
unsupported("torchmlx.optim.AdamW with non-default options") unsupported("torchmlx.optim.AdamW with non-default options")
model = getattr(params, "model", None)
if model is None:
raise TypeError(
"MLX AdamW requires parameters returned directly by model.parameters()"
)
self._parameters = params self._parameters = params
self._model = model
self._optimizer = _native.AdamW( self._optimizer = _native.AdamW(
learning_rate=lr, learning_rate=lr,
betas=list(betas), betas=list(betas),
@@ -37,16 +43,27 @@ else:
weight_decay=weight_decay, weight_decay=weight_decay,
bias_correction=True, bias_correction=True,
) )
self._pending_update = None
self._random_before_forward = None
@property @property
def state(self): def state(self):
return self._optimizer.state return self._optimizer.state
def zero_grad(self, *args, **kwargs): def zero_grad(self, *args, **kwargs):
unsupported("torchmlx.optim.AdamW.zero_grad on MLX; use torchmlx.Trainer") if args or kwargs:
unsupported("torchmlx.optim.AdamW.zero_grad with arguments")
from torchmlx._autograd import activate
self._pending_update = None
activate(self)
def step(self, *args, **kwargs): def step(self, *args, **kwargs):
unsupported("torchmlx.optim.AdamW.step on MLX; use torchmlx.Trainer") if args or kwargs:
unsupported("torchmlx.optim.AdamW.step with arguments")
from torchmlx._autograd import step
step(self)
def __getattr__(name): def __getattr__(name):
unsupported(f"torchmlx.optim.{name}") unsupported(f"torchmlx.optim.{name}")