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.
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.
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.
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)
model = TinyStoriesTransformer(config).to(device)
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):
input_tokens, target_tokens = create_batch(
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(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
_reshape = mx.array.reshape
_squeeze = mx.array.squeeze
_mean = mx.array.mean
_var = mx.array.var
_any = mx.array.any
_getitem = mx.array.__getitem__
_setitem = mx.array.__setitem__
_item = mx.array.item
def _torch_transpose(self, dim0=None, dim1=None):
@@ -18,17 +21,49 @@ def _torch_transpose(self, dim0=None, dim1=None):
else:
axes = list(range(self.ndim))
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):
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
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):
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):
@@ -126,9 +161,17 @@ def _requires_grad(self, requires_grad=True):
def _backward(self, *args, **kwargs):
raise RuntimeError(
"loss.backward() is not supported by the MLX backend; use torchmlx.Trainer"
)
if args or kwargs:
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):
@@ -144,9 +187,17 @@ def _torch_getitem(self, key):
if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key)
if key.shape == self.shape:
return _getitem(self.reshape(-1), indices)
return _getitem(self, indices)
return _getitem(self, key)
result = _getitem(self.reshape(-1), indices)
operation = lambda value: _getitem(_reshape(value, (-1,)), indices)
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):
@@ -163,8 +214,11 @@ def _torch_setitem(self, key, value):
def install(device_type):
mx.array.transpose = _torch_transpose
mx.array.reshape = _torch_reshape
mx.array.view = _view
mx.array.unsqueeze = _unsqueeze
mx.array.squeeze = _torch_squeeze
mx.array.flatten = _flatten
mx.array.float = _float
mx.array.bool = _bool
mx.array.pow = _pow
@@ -180,6 +234,7 @@ def install(device_type):
mx.array.contiguous = _contiguous
mx.array.requires_grad_ = _requires_grad
mx.array.backward = _backward
mx.array.item = _torch_item
mx.array.device = property(lambda self: device_type("mps"))
mx.array.__getitem__ = _torch_getitem
mx.array.__setitem__ = _torch_setitem
+18 -1
View File
@@ -23,15 +23,32 @@ else:
import mlx.core as mx
import mlx.nn as _native
class _ParameterTree(dict):
def __init__(self, values, model):
super().__init__(values)
self.model = model
class Module(_native.Module):
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):
raise NotImplementedError(
f"Module [{type(self).__name__}] is missing the required forward function"
)
def parameters(self):
return _ParameterTree(super().parameters(), self)
def to(self, *args, **kwargs):
dtype = kwargs.pop("dtype", None)
device = kwargs.pop("device", None)
+22 -3
View File
@@ -27,6 +27,25 @@ else:
reduction="mean",
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:
unsupported("torchmlx.nn.functional.cross_entropy legacy reductions")
if input.ndim < 2:
@@ -56,15 +75,15 @@ else:
losses = (1 - label_smoothing) * target_losses + label_smoothing * smooth_losses
losses = mx.where(valid, losses, mx.zeros_like(losses))
if reduction == "none":
return losses.reshape(output_shape)
return finish(losses.reshape(output_shape))
if reduction == "sum":
return mx.sum(losses)
return finish(mx.sum(losses))
if reduction == "mean":
if weight is None:
denominator = mx.sum(valid)
else:
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}")
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:
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._model = model
self._optimizer = _native.AdamW(
learning_rate=lr,
betas=list(betas),
@@ -37,16 +43,27 @@ else:
weight_decay=weight_decay,
bias_correction=True,
)
self._pending_update = None
self._random_before_forward = None
@property
def state(self):
return self._optimizer.state
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):
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):
unsupported(f"torchmlx.optim.{name}")