diff --git a/.gitignore b/.gitignore index dd0bded..9efc18b 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,4 @@ examples/tinystories-llm/TinyStoriesV2-GPT4-valid.txt # Environment files .env +AGENTS.md diff --git a/docs/compatibility.md b/docs/compatibility.md index 336c5f5..a3f1a29 100644 --- a/docs/compatibility.md +++ b/docs/compatibility.md @@ -18,7 +18,7 @@ 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. +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. Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation. diff --git a/src/torchmlx/_autograd.py b/src/torchmlx/_autograd.py index 24de777..2f07208 100644 --- a/src/torchmlx/_autograd.py +++ b/src/torchmlx/_autograd.py @@ -11,6 +11,78 @@ _replayed_losses = {} _suspended = False +class _ArraySlot: + def __init__(self, index): + self.index = index + + +class _BackwardPlan: + def __init__(self, value_and_grad, state): + self._value_and_grad = value_and_grad + self._compiled = mx.compile( + value_and_grad, + inputs=state, + outputs=state, + ) + self._compile_enabled = True + + def __call__(self, *inputs): + function = self._compiled if self._compile_enabled else self._value_and_grad + return function(*inputs) + + def fallback(self, *inputs): + self._compile_enabled = False + return self._value_and_grad(*inputs) + + +def _partition_arrays(value, arrays): + if isinstance(value, mx.array): + slot = _ArraySlot(len(arrays)) + arrays.append(value) + return slot + if isinstance(value, tuple): + return tuple(_partition_arrays(item, arrays) for item in value) + if isinstance(value, list): + return [_partition_arrays(item, arrays) for item in value] + if isinstance(value, dict): + return { + key: _partition_arrays(item, arrays) for key, item in value.items() + } + return value + + +def _restore_arrays(template, arrays): + if isinstance(template, _ArraySlot): + return arrays[template.index] + if isinstance(template, tuple): + return tuple(_restore_arrays(item, arrays) for item in template) + if isinstance(template, list): + return [_restore_arrays(item, arrays) for item in template] + if isinstance(template, dict): + return { + key: _restore_arrays(item, arrays) for key, item in template.items() + } + return template + + +def _template_key(template): + if isinstance(template, _ArraySlot): + return ("array",) + if isinstance(template, tuple): + return ("tuple", tuple(_template_key(item) for item in template)) + if isinstance(template, list): + return ("list", tuple(_template_key(item) for item in template)) + if isinstance(template, dict): + return ( + "dict", + tuple( + (type(key).__qualname__, repr(key), _template_key(item)) + for key, item in template.items() + ), + ) + return ("static", type(template).__qualname__, repr(template)) + + def _snapshot_random_state(): state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state] mx.eval(state) @@ -38,7 +110,7 @@ def end_forward(model, args, kwargs, output): and not _suspended and isinstance(output, mx.array) ): - _latest_forward = (model, args, kwargs, lambda value: value) + _latest_forward = (model, args, kwargs, ()) _lineage[id(output)] = (output, _latest_forward) @@ -57,37 +129,69 @@ def activate(optimizer): optimizer._random_before_forward = _snapshot_random_state() -def propagate(source, result, operation=lambda value: value): +def propagate(source, result, operation, signature, operands=()): entry = _lineage.get(id(source)) if not _suspended and entry is not None and entry[0] is source: - model, args, kwargs, previous = entry[1] + model, args, kwargs, transforms = entry[1] _lineage[id(result)] = ( result, ( model, args, kwargs, - lambda output: operation(previous(output)), + transforms + ((operation, signature, operands),), ), ) return result -def register_loss(loss, loss_input, rebuild): +def register_loss(loss, loss_input, rebuild, operands, signature): 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] + model, args, kwargs, transforms = entry[1] - def plan(): - def objective(): - return rebuild(transform(model(*args, **kwargs))) + context = ( + args, + kwargs, + tuple(transform_operands for _, _, transform_operands in transforms), + operands, + ) + dynamic_inputs = [] + template = _partition_arrays(context, dynamic_inputs) + cache_key = ( + getattr(model, "training", None), + _template_key(template), + tuple(transform_signature for _, transform_signature, _ in transforms), + signature, + ) - return nn.value_and_grad(model, objective)() + def build(): + state = [model.state, mx.random.state] - _loss_plans[id(loss)] = (loss, model, 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) + + value_and_grad = nn.value_and_grad(model, objective) + return _BackwardPlan(value_and_grad, state) + + _loss_plans[id(loss)] = ( + loss, + model, + cache_key, + tuple(dynamic_inputs), + build, + ) return loss @@ -100,12 +204,32 @@ def backward(loss): raise RuntimeError( "this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function" ) - _, model, plan = entry + _, model, cache_key, dynamic_inputs, build = entry random_after_forward = _snapshot_random_state() _restore_random_state(_active_optimizer._random_before_forward) _suspended = True try: - replayed_loss, gradients = plan() + 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: + replayed_loss, gradients = plan(*dynamic_inputs) + except ValueError as error: + message = str(error) + if ( + not plan._compile_enabled + or "Attempting to eval an array" not in message + or not ( + "function transformations" in message + or "without a primitive" in message + ) + ): + 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: _suspended = False @@ -126,8 +250,8 @@ def step(optimizer): 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._compiled_step(gradients) + mx.eval(optimizer._step_state) optimizer._pending_update = None _active_optimizer = None _latest_forward = None diff --git a/src/torchmlx/_mlx_tensor.py b/src/torchmlx/_mlx_tensor.py index 7000257..55aeb9d 100644 --- a/src/torchmlx/_mlx_tensor.py +++ b/src/torchmlx/_mlx_tensor.py @@ -24,7 +24,12 @@ def _torch_transpose(self, dim0=None, dim1=None): result = _transpose(self, axes) from ._autograd import propagate - return propagate(self, result, lambda value: _transpose(value, axes)) + return propagate( + self, + result, + lambda value: _transpose(value, axes), + ("transpose", tuple(axes)), + ) def _view(self, *shape): @@ -39,21 +44,36 @@ def _torch_reshape(self, *shape): result = _reshape(self, shape) from ._autograd import propagate - return propagate(self, result, lambda value: _reshape(value, shape)) + return propagate( + self, + result, + lambda value: _reshape(value, shape), + ("reshape", tuple(shape)), + ) def _unsqueeze(self, dim): result = mx.expand_dims(self, axis=dim) from ._autograd import propagate - return propagate(self, result, lambda value: mx.expand_dims(value, axis=dim)) + return propagate( + self, + result, + lambda value: mx.expand_dims(value, axis=dim), + ("unsqueeze", 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)) + return propagate( + self, + result, + lambda value: _squeeze(value, axis=dim), + ("squeeze", dim), + ) def _flatten(self, start_dim=0, end_dim=-1): @@ -188,16 +208,30 @@ def _torch_getitem(self, key): indices = _mask_indices(key) if key.shape == self.shape: result = _getitem(self.reshape(-1), indices) - operation = lambda value: _getitem(_reshape(value, (-1,)), indices) + operation = lambda value, current_indices: _getitem( + _reshape(value, (-1,)), current_indices + ) + signature = ("getitem_bool", "flat") else: result = _getitem(self, indices) - operation = lambda value: _getitem(value, indices) + operation = lambda value, current_indices: _getitem( + value, current_indices + ) + signature = ("getitem_bool", "first_axis") + operands = (indices,) + elif isinstance(key, mx.array): + result = _getitem(self, key) + operation = lambda value, current_key: _getitem(value, current_key) + signature = ("getitem_array",) + operands = (key,) else: result = _getitem(self, key) operation = lambda value: _getitem(value, key) + signature = ("getitem", type(key).__qualname__, repr(key)) + operands = () from ._autograd import propagate - return propagate(self, result, operation) + return propagate(self, result, operation, signature, operands) def _torch_setitem(self, key, value): diff --git a/src/torchmlx/nn/functional.py b/src/torchmlx/nn/functional.py index cbda73e..6eaa5fd 100644 --- a/src/torchmlx/nn/functional.py +++ b/src/torchmlx/nn/functional.py @@ -32,11 +32,11 @@ else: def finish(loss): from torchmlx._autograd import register_loss - def rebuild(recomputed_input): + def rebuild(recomputed_input, current_target, *current_weight): return cross_entropy( recomputed_input, - target, - weight=weight, + current_target, + weight=current_weight[0] if current_weight else None, size_average=size_average, ignore_index=ignore_index, reduce=reduce, @@ -44,7 +44,23 @@ else: label_smoothing=label_smoothing, ) - return register_loss(loss, original_input, rebuild) + operands = (target,) if weight is None else (target, weight) + signature = ( + "cross_entropy", + weight is not None, + size_average, + ignore_index, + reduce, + reduction, + label_smoothing, + ) + return register_loss( + loss, + original_input, + rebuild, + operands, + signature, + ) if size_average is not None or reduce is not None: unsupported("torchmlx.nn.functional.cross_entropy legacy reductions") diff --git a/src/torchmlx/optim/__init__.py b/src/torchmlx/optim/__init__.py index 7cc10cf..7653ff3 100644 --- a/src/torchmlx/optim/__init__.py +++ b/src/torchmlx/optim/__init__.py @@ -10,6 +10,7 @@ if BACKEND == "torch": return getattr(_native, name) else: + import mlx.core as mx import mlx.optimizers as _native class AdamW: @@ -45,6 +46,19 @@ else: ) self._pending_update = None self._random_before_forward = None + self._compiled_backward = {} + self._optimizer.init(model.trainable_parameters()) + state = [model.state, self._optimizer.state, mx.random.state] + + def update(gradients): + self._optimizer.update(model, gradients) + + self._compiled_step = mx.compile( + update, + inputs=state, + outputs=state, + ) + self._step_state = state @property def state(self):