mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
support backward training
This commit is contained in:
1 parent
7fa8ae82ad
commit
53c9cf4159
7 files changed
+277
-18
No files matched your search
+13
-1
@@ -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.
|
||||
@@ -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))
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in new issue
Block a user