diff --git a/docs/compatibility.md b/docs/compatibility.md index 7126ed4..336c5f5 100644 --- a/docs/compatibility.md +++ b/docs/compatibility.md @@ -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. diff --git a/examples/tinystories-llm/train.py b/examples/tinystories-llm/train.py index 870b827..2f48a5e 100644 --- a/examples/tinystories-llm/train.py +++ b/examples/tinystories-llm/train.py @@ -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)) diff --git a/src/torchmlx/_autograd.py b/src/torchmlx/_autograd.py new file mode 100644 index 0000000..24de777 --- /dev/null +++ b/src/torchmlx/_autograd.py @@ -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() diff --git a/src/torchmlx/_mlx_tensor.py b/src/torchmlx/_mlx_tensor.py index 7057def..7000257 100644 --- a/src/torchmlx/_mlx_tensor.py +++ b/src/torchmlx/_mlx_tensor.py @@ -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 diff --git a/src/torchmlx/nn/__init__.py b/src/torchmlx/nn/__init__.py index 2715f19..b3af388 100644 --- a/src/torchmlx/nn/__init__.py +++ b/src/torchmlx/nn/__init__.py @@ -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) diff --git a/src/torchmlx/nn/functional.py b/src/torchmlx/nn/functional.py index a9d0931..cbda73e 100644 --- a/src/torchmlx/nn/functional.py +++ b/src/torchmlx/nn/functional.py @@ -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( diff --git a/src/torchmlx/optim/__init__.py b/src/torchmlx/optim/__init__.py index 74c241a..7cc10cf 100644 --- a/src/torchmlx/optim/__init__.py +++ b/src/torchmlx/optim/__init__.py @@ -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}")