mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
fuse mlx training step
This commit is contained in:
1 parent
69489225c3
commit
4eb86111a8
6 files changed
+352
-124
No files matched your search
@@ -18,7 +18,7 @@ loss.backward()
|
|||||||
optimizer.step()
|
optimizer.step()
|
||||||
```
|
```
|
||||||
|
|
||||||
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.
|
MLX implements this sequence by recording the outer model call and replaying the forward, backward, and AdamW update inside one cached compiled graph during `step`. 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. Reading `loss.item()` before `step` materializes gradients through a separate compiled path. 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.
|
Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation.
|
||||||
|
|
||||||
|
|||||||
+341
-91
@@ -4,7 +4,6 @@ import mlx.nn as nn
|
|||||||
|
|
||||||
_active_optimizer = None
|
_active_optimizer = None
|
||||||
_forward_depth = 0
|
_forward_depth = 0
|
||||||
_latest_forward = None
|
|
||||||
_loss_plans = {}
|
_loss_plans = {}
|
||||||
_lineage = {}
|
_lineage = {}
|
||||||
_replayed_losses = {}
|
_replayed_losses = {}
|
||||||
@@ -16,24 +15,48 @@ class _ArraySlot:
|
|||||||
self.index = index
|
self.index = index
|
||||||
|
|
||||||
|
|
||||||
class _BackwardPlan:
|
class _TrainingPlan:
|
||||||
def __init__(self, value_and_grad, state):
|
def __init__(self, model, optimizer, objective, compile=True):
|
||||||
self._value_and_grad = value_and_grad
|
self.model = model
|
||||||
self._compiled = mx.compile(
|
self.optimizer = optimizer
|
||||||
value_and_grad,
|
self._value_and_grad = nn.value_and_grad(model, objective)
|
||||||
inputs=state,
|
self._compile_enabled = compile
|
||||||
outputs=state,
|
self._compiled_backward = None
|
||||||
|
self._compiled_fused = None
|
||||||
|
self._backward_signatures = set()
|
||||||
|
self._fused_signatures = set()
|
||||||
|
if compile:
|
||||||
|
backward_state = [model.state, mx.random.state]
|
||||||
|
fused_state = [model.state, optimizer.state, mx.random.state]
|
||||||
|
self._compiled_backward = mx.compile(
|
||||||
|
self._value_and_grad,
|
||||||
|
inputs=backward_state,
|
||||||
|
outputs=backward_state,
|
||||||
|
)
|
||||||
|
self._compiled_fused = mx.compile(
|
||||||
|
self._fused,
|
||||||
|
inputs=fused_state,
|
||||||
|
outputs=fused_state,
|
||||||
)
|
)
|
||||||
self._compile_enabled = True
|
|
||||||
|
|
||||||
def __call__(self, *inputs):
|
def _fused(self, *inputs):
|
||||||
function = self._compiled if self._compile_enabled else self._value_and_grad
|
loss, gradients = self._value_and_grad(*inputs)
|
||||||
return function(*inputs)
|
self.optimizer.update(self.model, gradients)
|
||||||
|
return loss
|
||||||
|
|
||||||
def fallback(self, *inputs):
|
def backward(self, *inputs):
|
||||||
self._compile_enabled = False
|
if self._compile_enabled:
|
||||||
|
return self._compiled_backward(*inputs)
|
||||||
return self._value_and_grad(*inputs)
|
return self._value_and_grad(*inputs)
|
||||||
|
|
||||||
|
def fused(self, *inputs):
|
||||||
|
if self._compile_enabled:
|
||||||
|
return self._compiled_fused(*inputs)
|
||||||
|
return self._fused(*inputs)
|
||||||
|
|
||||||
|
def disable_compile(self):
|
||||||
|
self._compile_enabled = False
|
||||||
|
|
||||||
|
|
||||||
def _partition_arrays(value, arrays):
|
def _partition_arrays(value, arrays):
|
||||||
if isinstance(value, mx.array):
|
if isinstance(value, mx.array):
|
||||||
@@ -45,9 +68,7 @@ def _partition_arrays(value, arrays):
|
|||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
return [_partition_arrays(item, arrays) for item in value]
|
return [_partition_arrays(item, arrays) for item in value]
|
||||||
if isinstance(value, dict):
|
if isinstance(value, dict):
|
||||||
return {
|
return {key: _partition_arrays(item, arrays) for key, item in value.items()}
|
||||||
key: _partition_arrays(item, arrays) for key, item in value.items()
|
|
||||||
}
|
|
||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
@@ -59,9 +80,7 @@ def _restore_arrays(template, arrays):
|
|||||||
if isinstance(template, list):
|
if isinstance(template, list):
|
||||||
return [_restore_arrays(item, arrays) for item in template]
|
return [_restore_arrays(item, arrays) for item in template]
|
||||||
if isinstance(template, dict):
|
if isinstance(template, dict):
|
||||||
return {
|
return {key: _restore_arrays(item, arrays) for key, item in template.items()}
|
||||||
key: _restore_arrays(item, arrays) for key, item in template.items()
|
|
||||||
}
|
|
||||||
return template
|
return template
|
||||||
|
|
||||||
|
|
||||||
@@ -83,6 +102,16 @@ def _template_key(template):
|
|||||||
return ("static", type(template).__qualname__, repr(template))
|
return ("static", type(template).__qualname__, repr(template))
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_tree(value):
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {key: _copy_tree(item) for key, item in value.items()}
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [_copy_tree(item) for item in value]
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return tuple(_copy_tree(item) for item in value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
def _snapshot_random_state():
|
def _snapshot_random_state():
|
||||||
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
|
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
|
||||||
mx.eval(state)
|
mx.eval(state)
|
||||||
@@ -95,23 +124,54 @@ def _restore_random_state(state):
|
|||||||
mx.eval(mx.random.state)
|
mx.eval(mx.random.state)
|
||||||
|
|
||||||
|
|
||||||
def begin_forward():
|
def _compile_failure(error):
|
||||||
|
message = str(error)
|
||||||
|
return "Attempting to eval an array" in message and (
|
||||||
|
"function transformations" in message or "without a primitive" in message
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def begin_forward(model):
|
||||||
global _forward_depth
|
global _forward_depth
|
||||||
|
if (
|
||||||
|
_suspended
|
||||||
|
or _active_optimizer is None
|
||||||
|
or model is not _active_optimizer._model
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
if _forward_depth:
|
||||||
|
_active_optimizer._recursive_forward = True
|
||||||
_forward_depth += 1
|
_forward_depth += 1
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def note_eager_evaluation():
|
||||||
|
if _forward_depth and _active_optimizer is not None and not _suspended:
|
||||||
|
_active_optimizer._eager_forward = True
|
||||||
|
|
||||||
|
|
||||||
def end_forward(model, args, kwargs, output):
|
def end_forward(model, args, kwargs, output):
|
||||||
global _forward_depth, _latest_forward
|
global _forward_depth
|
||||||
_forward_depth -= 1
|
_forward_depth -= 1
|
||||||
if (
|
if _forward_depth != 0 or not isinstance(output, mx.array):
|
||||||
_forward_depth == 0
|
return
|
||||||
and _active_optimizer is not None
|
random_before = _active_optimizer._random_before_forward
|
||||||
and model is _active_optimizer._model
|
stochastic = any(
|
||||||
and not _suspended
|
current is not previous
|
||||||
and isinstance(output, mx.array)
|
for current, previous in zip(mx.random.state, random_before)
|
||||||
):
|
)
|
||||||
_latest_forward = (model, args, kwargs, ())
|
descriptor = (
|
||||||
_lineage[id(output)] = (output, _latest_forward)
|
model,
|
||||||
|
args,
|
||||||
|
kwargs,
|
||||||
|
(),
|
||||||
|
random_before if stochastic else None,
|
||||||
|
not (
|
||||||
|
_active_optimizer._recursive_forward
|
||||||
|
or _active_optimizer._eager_forward
|
||||||
|
),
|
||||||
|
)
|
||||||
|
_lineage[id(output)] = (output, descriptor)
|
||||||
|
|
||||||
|
|
||||||
def abort_forward():
|
def abort_forward():
|
||||||
@@ -120,19 +180,24 @@ def abort_forward():
|
|||||||
|
|
||||||
|
|
||||||
def activate(optimizer):
|
def activate(optimizer):
|
||||||
global _active_optimizer, _latest_forward
|
global _active_optimizer, _forward_depth
|
||||||
_active_optimizer = optimizer
|
_active_optimizer = optimizer
|
||||||
_latest_forward = None
|
_forward_depth = 0
|
||||||
_loss_plans.clear()
|
_loss_plans.clear()
|
||||||
_lineage.clear()
|
_lineage.clear()
|
||||||
_replayed_losses.clear()
|
_replayed_losses.clear()
|
||||||
optimizer._random_before_forward = _snapshot_random_state()
|
optimizer._random_before_forward = tuple(mx.random.state)
|
||||||
|
optimizer._recursive_forward = False
|
||||||
|
optimizer._eager_forward = False
|
||||||
|
|
||||||
|
|
||||||
def propagate(source, result, operation, signature, operands=()):
|
def propagate(source, result, operation, signature, operands=()):
|
||||||
|
if _suspended:
|
||||||
|
return result
|
||||||
entry = _lineage.get(id(source))
|
entry = _lineage.get(id(source))
|
||||||
if not _suspended and entry is not None and entry[0] is source:
|
if entry is None or entry[0] is not source:
|
||||||
model, args, kwargs, transforms = entry[1]
|
return result
|
||||||
|
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
|
||||||
_lineage[id(result)] = (
|
_lineage[id(result)] = (
|
||||||
result,
|
result,
|
||||||
(
|
(
|
||||||
@@ -140,19 +205,45 @@ def propagate(source, result, operation, signature, operands=()):
|
|||||||
args,
|
args,
|
||||||
kwargs,
|
kwargs,
|
||||||
transforms + ((operation, signature, operands),),
|
transforms + ((operation, signature, operands),),
|
||||||
|
random_before,
|
||||||
|
compile_allowed,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _plan_for(
|
||||||
|
optimizer, model, cache_key, template, transforms, rebuild, compile_allowed
|
||||||
|
):
|
||||||
|
plan = optimizer._training_plans.get(cache_key)
|
||||||
|
if plan is not None:
|
||||||
|
return 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)
|
||||||
|
|
||||||
|
plan = _TrainingPlan(
|
||||||
|
model, optimizer._optimizer, objective, compile=compile_allowed
|
||||||
|
)
|
||||||
|
optimizer._training_plans[cache_key] = plan
|
||||||
|
return plan
|
||||||
|
|
||||||
|
|
||||||
def register_loss(loss, loss_input, rebuild, operands, signature):
|
def register_loss(loss, loss_input, rebuild, operands, signature):
|
||||||
if _suspended or _active_optimizer is None:
|
if _suspended or _active_optimizer is None:
|
||||||
return loss
|
return loss
|
||||||
entry = _lineage.get(id(loss_input))
|
entry = _lineage.get(id(loss_input))
|
||||||
if entry is None or entry[0] is not loss_input:
|
if entry is None or entry[0] is not loss_input:
|
||||||
return loss
|
return loss
|
||||||
model, args, kwargs, transforms = entry[1]
|
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
|
||||||
|
|
||||||
context = (
|
context = (
|
||||||
args,
|
args,
|
||||||
kwargs,
|
kwargs,
|
||||||
@@ -166,37 +257,23 @@ def register_loss(loss, loss_input, rebuild, operands, signature):
|
|||||||
_template_key(template),
|
_template_key(template),
|
||||||
tuple(transform_signature for _, transform_signature, _ in transforms),
|
tuple(transform_signature for _, transform_signature, _ in transforms),
|
||||||
signature,
|
signature,
|
||||||
|
compile_allowed,
|
||||||
)
|
)
|
||||||
|
|
||||||
def build():
|
|
||||||
state = [model.state, mx.random.state]
|
|
||||||
|
|
||||||
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_plans[id(loss)] = (
|
||||||
loss,
|
loss,
|
||||||
model,
|
model,
|
||||||
cache_key,
|
cache_key,
|
||||||
tuple(dynamic_inputs),
|
tuple(dynamic_inputs),
|
||||||
build,
|
template,
|
||||||
|
transforms,
|
||||||
|
rebuild,
|
||||||
|
random_before,
|
||||||
|
compile_allowed,
|
||||||
)
|
)
|
||||||
return loss
|
return loss
|
||||||
|
|
||||||
|
|
||||||
def backward(loss):
|
def backward(loss):
|
||||||
global _suspended
|
|
||||||
if _active_optimizer is None:
|
if _active_optimizer is None:
|
||||||
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
|
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
|
||||||
entry = _loss_plans.get(id(loss))
|
entry = _loss_plans.get(id(loss))
|
||||||
@@ -204,56 +281,229 @@ def backward(loss):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
|
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
|
||||||
)
|
)
|
||||||
_, model, cache_key, dynamic_inputs, build = entry
|
(
|
||||||
random_after_forward = _snapshot_random_state()
|
_,
|
||||||
_restore_random_state(_active_optimizer._random_before_forward)
|
model,
|
||||||
|
cache_key,
|
||||||
|
dynamic_inputs,
|
||||||
|
template,
|
||||||
|
transforms,
|
||||||
|
rebuild,
|
||||||
|
random_before,
|
||||||
|
compile_allowed,
|
||||||
|
) = entry
|
||||||
|
plan = _plan_for(
|
||||||
|
_active_optimizer,
|
||||||
|
model,
|
||||||
|
cache_key,
|
||||||
|
template,
|
||||||
|
transforms,
|
||||||
|
rebuild,
|
||||||
|
compile_allowed,
|
||||||
|
)
|
||||||
|
_active_optimizer._pending_update = (
|
||||||
|
"staged",
|
||||||
|
loss,
|
||||||
|
plan,
|
||||||
|
dynamic_inputs,
|
||||||
|
random_before,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _restore_execution(model, optimizer, parameters, optimizer_state):
|
||||||
|
model.update(parameters)
|
||||||
|
if optimizer_state is not None:
|
||||||
|
_restore_tree(optimizer.state, optimizer_state)
|
||||||
|
|
||||||
|
|
||||||
|
def _restore_tree(current, saved):
|
||||||
|
if isinstance(current, dict) and isinstance(saved, dict):
|
||||||
|
for key in tuple(current):
|
||||||
|
if key not in saved:
|
||||||
|
del current[key]
|
||||||
|
for key, value in saved.items():
|
||||||
|
if key in current:
|
||||||
|
restored = _restore_tree(current[key], value)
|
||||||
|
if restored is not current[key]:
|
||||||
|
current[key] = restored
|
||||||
|
else:
|
||||||
|
current[key] = value
|
||||||
|
return current
|
||||||
|
if isinstance(current, list) and isinstance(saved, list):
|
||||||
|
restored = [
|
||||||
|
_restore_tree(old, value) for old, value in zip(current, saved)
|
||||||
|
]
|
||||||
|
current[:] = restored + saved[len(restored) :]
|
||||||
|
return current
|
||||||
|
return saved
|
||||||
|
|
||||||
|
|
||||||
|
def _execute(plan, inputs, fused, random_before, advance_random=False):
|
||||||
|
global _suspended
|
||||||
|
model = plan.model
|
||||||
|
optimizer = plan.optimizer
|
||||||
|
signatures = plan._fused_signatures if fused else plan._backward_signatures
|
||||||
|
signature = tuple((value.shape, value.dtype) for value in inputs)
|
||||||
|
protect_state = not plan._compile_enabled or signature not in signatures
|
||||||
|
parameters = (
|
||||||
|
_copy_tree(model.trainable_parameters())
|
||||||
|
if protect_state
|
||||||
|
else getattr(optimizer, "_torchmlx_parameters", None)
|
||||||
|
)
|
||||||
|
optimizer_state = None
|
||||||
|
if fused:
|
||||||
|
optimizer_state = (
|
||||||
|
_copy_tree(optimizer.state)
|
||||||
|
if protect_state
|
||||||
|
else getattr(optimizer, "_torchmlx_state", None)
|
||||||
|
)
|
||||||
|
random_after = _snapshot_random_state() if random_before is not None else None
|
||||||
|
random_rollback = tuple(mx.random.state) if advance_random else None
|
||||||
|
if random_before is not None:
|
||||||
|
_restore_random_state(random_before)
|
||||||
_suspended = True
|
_suspended = True
|
||||||
try:
|
try:
|
||||||
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:
|
try:
|
||||||
replayed_loss, gradients = plan(*dynamic_inputs)
|
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
|
||||||
except ValueError as error:
|
current_parameters = model.parameters()
|
||||||
message = str(error)
|
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
|
||||||
if (
|
mx.eval(
|
||||||
not plan._compile_enabled
|
result,
|
||||||
or "Attempting to eval an array" not in message
|
current_parameters,
|
||||||
or not (
|
current_optimizer_state,
|
||||||
"function transformations" in message
|
mx.random.state
|
||||||
or "without a primitive" in message
|
if random_before is not None or advance_random
|
||||||
|
else [],
|
||||||
)
|
)
|
||||||
):
|
signatures.add(signature)
|
||||||
|
if fused:
|
||||||
|
optimizer._torchmlx_parameters = current_parameters
|
||||||
|
optimizer._torchmlx_state = current_optimizer_state
|
||||||
|
except (ValueError, RuntimeError) as error:
|
||||||
|
if not plan._compile_enabled or not _compile_failure(error):
|
||||||
|
raise
|
||||||
|
_restore_execution(model, optimizer, parameters, optimizer_state)
|
||||||
|
if random_before is not None:
|
||||||
|
_restore_random_state(random_before)
|
||||||
|
elif random_rollback is not None:
|
||||||
|
_restore_random_state(random_rollback)
|
||||||
|
plan.disable_compile()
|
||||||
|
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
|
||||||
|
current_parameters = model.parameters()
|
||||||
|
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
|
||||||
|
mx.eval(
|
||||||
|
result,
|
||||||
|
current_parameters,
|
||||||
|
current_optimizer_state,
|
||||||
|
mx.random.state
|
||||||
|
if random_before is not None or advance_random
|
||||||
|
else [],
|
||||||
|
)
|
||||||
|
if fused:
|
||||||
|
optimizer._torchmlx_parameters = current_parameters
|
||||||
|
optimizer._torchmlx_state = current_optimizer_state
|
||||||
|
except Exception:
|
||||||
|
if parameters is not None:
|
||||||
|
_restore_execution(model, optimizer, parameters, optimizer_state)
|
||||||
|
if random_rollback is not None:
|
||||||
|
_restore_random_state(random_rollback)
|
||||||
raise
|
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:
|
finally:
|
||||||
_suspended = False
|
_suspended = False
|
||||||
_restore_random_state(random_after_forward)
|
if random_after is not None:
|
||||||
|
_restore_random_state(random_after)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _materialize(optimizer):
|
||||||
|
pending = optimizer._pending_update
|
||||||
|
_, loss, plan, dynamic_inputs, random_before = pending
|
||||||
|
replayed_loss, gradients = _execute(
|
||||||
|
plan, dynamic_inputs, fused=False, random_before=random_before
|
||||||
|
)
|
||||||
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
||||||
_active_optimizer._pending_update = (model, gradients)
|
optimizer._pending_update = ("materialized", loss, plan.model, gradients)
|
||||||
|
return replayed_loss
|
||||||
|
|
||||||
|
|
||||||
def replayed_value(value):
|
def replayed_value(value):
|
||||||
|
if _active_optimizer is not None:
|
||||||
|
pending = _active_optimizer._pending_update
|
||||||
|
if pending is not None and pending[0] == "staged" and pending[1] is value:
|
||||||
|
return _materialize(_active_optimizer)
|
||||||
entry = _replayed_losses.get(id(value))
|
entry = _replayed_losses.get(id(value))
|
||||||
if entry is None or entry[0] is not value:
|
if entry is None or entry[0] is not value:
|
||||||
return value
|
return value
|
||||||
return entry[1]
|
return entry[1]
|
||||||
|
|
||||||
|
|
||||||
def step(optimizer):
|
def _finish_step(loss, replayed_loss):
|
||||||
global _active_optimizer, _latest_forward
|
global _active_optimizer, _forward_depth
|
||||||
if optimizer._pending_update is None:
|
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
||||||
raise RuntimeError("loss.backward() must be called before optimizer.step()")
|
_active_optimizer._pending_update = None
|
||||||
model, gradients = optimizer._pending_update
|
|
||||||
optimizer._compiled_step(gradients)
|
|
||||||
mx.eval(optimizer._step_state)
|
|
||||||
optimizer._pending_update = None
|
|
||||||
_active_optimizer = None
|
_active_optimizer = None
|
||||||
_latest_forward = None
|
_forward_depth = 0
|
||||||
_loss_plans.clear()
|
_loss_plans.clear()
|
||||||
_lineage.clear()
|
_lineage.clear()
|
||||||
|
|
||||||
|
|
||||||
|
def step(optimizer):
|
||||||
|
pending = optimizer._pending_update
|
||||||
|
if pending is None:
|
||||||
|
raise RuntimeError("loss.backward() must be called before optimizer.step()")
|
||||||
|
if pending[0] == "staged":
|
||||||
|
_, loss, plan, dynamic_inputs, random_before = pending
|
||||||
|
replayed_loss = _execute(
|
||||||
|
plan, dynamic_inputs, fused=True, random_before=random_before
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
_, loss, model, gradients = pending
|
||||||
|
parameters = _copy_tree(model.trainable_parameters())
|
||||||
|
optimizer_state = _copy_tree(optimizer._optimizer.state)
|
||||||
|
try:
|
||||||
|
optimizer._compiled_step(gradients)
|
||||||
|
current_parameters = model.parameters()
|
||||||
|
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
|
||||||
|
mx.eval(current_parameters, current_optimizer_state)
|
||||||
|
except Exception:
|
||||||
|
_restore_execution(
|
||||||
|
model, optimizer._optimizer, parameters, optimizer_state
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
optimizer._optimizer._torchmlx_parameters = current_parameters
|
||||||
|
optimizer._optimizer._torchmlx_state = current_optimizer_state
|
||||||
|
replayed_loss = _replayed_losses[id(loss)][1]
|
||||||
|
_finish_step(loss, replayed_loss)
|
||||||
|
|
||||||
|
|
||||||
|
def trainer_step(trainer, x, y):
|
||||||
|
context = (x, y)
|
||||||
|
dynamic_inputs = []
|
||||||
|
template = _partition_arrays(context, dynamic_inputs)
|
||||||
|
cache_key = (
|
||||||
|
getattr(trainer.model, "training", None),
|
||||||
|
_template_key(template),
|
||||||
|
(),
|
||||||
|
("trainer", id(trainer.loss_fn), trainer.compile),
|
||||||
|
)
|
||||||
|
plan = trainer.optimizer._training_plans.get(cache_key)
|
||||||
|
if plan is None:
|
||||||
|
|
||||||
|
def objective(*current_inputs):
|
||||||
|
current_x, current_y = _restore_arrays(template, current_inputs)
|
||||||
|
return trainer.loss_fn(trainer.model(current_x), current_y)
|
||||||
|
|
||||||
|
plan = _TrainingPlan(
|
||||||
|
trainer.model,
|
||||||
|
trainer.optimizer._optimizer,
|
||||||
|
objective,
|
||||||
|
compile=trainer.compile,
|
||||||
|
)
|
||||||
|
trainer.optimizer._training_plans[cache_key] = plan
|
||||||
|
return _execute(
|
||||||
|
plan,
|
||||||
|
tuple(dynamic_inputs),
|
||||||
|
fused=True,
|
||||||
|
random_before=None,
|
||||||
|
advance_random=True,
|
||||||
|
)
|
||||||
@@ -189,8 +189,9 @@ def _backward(self, *args, **kwargs):
|
|||||||
|
|
||||||
|
|
||||||
def _torch_item(self):
|
def _torch_item(self):
|
||||||
from ._autograd import replayed_value
|
from ._autograd import note_eager_evaluation, replayed_value
|
||||||
|
|
||||||
|
note_eager_evaluation()
|
||||||
return _item(replayed_value(self))
|
return _item(replayed_value(self))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -32,12 +32,14 @@ else:
|
|||||||
def __call__(self, *args, **kwargs):
|
def __call__(self, *args, **kwargs):
|
||||||
from torchmlx._autograd import abort_forward, begin_forward, end_forward
|
from torchmlx._autograd import abort_forward, begin_forward, end_forward
|
||||||
|
|
||||||
begin_forward()
|
tracking = begin_forward(self)
|
||||||
try:
|
try:
|
||||||
output = self.forward(*args, **kwargs)
|
output = self.forward(*args, **kwargs)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
if tracking:
|
||||||
abort_forward()
|
abort_forward()
|
||||||
raise
|
raise
|
||||||
|
if tracking:
|
||||||
end_forward(self, args, kwargs, output)
|
end_forward(self, args, kwargs, output)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ else:
|
|||||||
)
|
)
|
||||||
self._pending_update = None
|
self._pending_update = None
|
||||||
self._random_before_forward = None
|
self._random_before_forward = None
|
||||||
self._compiled_backward = {}
|
self._training_plans = {}
|
||||||
self._optimizer.init(model.trainable_parameters())
|
self._optimizer.init(model.trainable_parameters())
|
||||||
state = [model.state, self._optimizer.state, mx.random.state]
|
state = [model.state, self._optimizer.state, mx.random.state]
|
||||||
|
|
||||||
|
|||||||
+4
-29
@@ -22,39 +22,14 @@ if BACKEND == "torch":
|
|||||||
return self._step(x, y)
|
return self._step(x, y)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
import mlx.core as mx
|
|
||||||
import mlx.nn as nn
|
|
||||||
|
|
||||||
class Trainer:
|
class Trainer:
|
||||||
def __init__(self, model, optimizer, loss_fn, compile=True):
|
def __init__(self, model, optimizer, loss_fn, compile=True):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.optimizer = optimizer
|
self.optimizer = optimizer
|
||||||
self.loss_fn = loss_fn
|
self.loss_fn = loss_fn
|
||||||
native_optimizer = optimizer._optimizer
|
self.compile = compile
|
||||||
native_optimizer.init(model.trainable_parameters())
|
|
||||||
|
|
||||||
def objective(x, y):
|
|
||||||
return loss_fn(model(x), y)
|
|
||||||
|
|
||||||
value_and_grad = nn.value_and_grad(model, objective)
|
|
||||||
|
|
||||||
def train_step(x, y):
|
|
||||||
loss, gradients = value_and_grad(x, y)
|
|
||||||
native_optimizer.update(model, gradients)
|
|
||||||
return loss
|
|
||||||
|
|
||||||
if compile:
|
|
||||||
state = [model.state, native_optimizer.state, mx.random.state]
|
|
||||||
self._step = mx.compile(train_step, inputs=state, outputs=state)
|
|
||||||
else:
|
|
||||||
self._step = train_step
|
|
||||||
|
|
||||||
def step(self, x, y):
|
def step(self, x, y):
|
||||||
loss = self._step(x, y)
|
from ._autograd import trainer_step
|
||||||
mx.eval(
|
|
||||||
loss,
|
return trainer_step(self, x, y)
|
||||||
self.model.parameters(),
|
|
||||||
self.optimizer._optimizer.state,
|
|
||||||
mx.random.state,
|
|
||||||
)
|
|
||||||
return loss
|
|
||||||
Reference in new issue
Block a user