mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 19:17: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.
|
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.
|
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.
|
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.
|
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.
|
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)
|
encoded_text = tokenizer.encode(training_text)
|
||||||
model = TinyStoriesTransformer(config).to(device)
|
model = TinyStoriesTransformer(config).to(device)
|
||||||
optimizer = optim.AdamW(model.parameters(), lr=3e-4)
|
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):
|
for step in range(arguments.steps):
|
||||||
input_tokens, target_tokens = create_batch(
|
input_tokens, target_tokens = create_batch(
|
||||||
encoded_text, batch_size=8, context_length=config.context_length
|
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(f"step {step + 1}: loss {loss.item():.4f}")
|
||||||
|
|
||||||
print(generate_text(model, tokenizer, "Once upon a time", token_count=120))
|
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
|
_transpose = mx.array.transpose
|
||||||
|
_reshape = mx.array.reshape
|
||||||
|
_squeeze = mx.array.squeeze
|
||||||
_mean = mx.array.mean
|
_mean = mx.array.mean
|
||||||
_var = mx.array.var
|
_var = mx.array.var
|
||||||
_any = mx.array.any
|
_any = mx.array.any
|
||||||
_getitem = mx.array.__getitem__
|
_getitem = mx.array.__getitem__
|
||||||
_setitem = mx.array.__setitem__
|
_setitem = mx.array.__setitem__
|
||||||
|
_item = mx.array.item
|
||||||
|
|
||||||
|
|
||||||
def _torch_transpose(self, dim0=None, dim1=None):
|
def _torch_transpose(self, dim0=None, dim1=None):
|
||||||
@@ -18,17 +21,49 @@ def _torch_transpose(self, dim0=None, dim1=None):
|
|||||||
else:
|
else:
|
||||||
axes = list(range(self.ndim))
|
axes = list(range(self.ndim))
|
||||||
axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
|
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):
|
def _view(self, *shape):
|
||||||
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
||||||
shape = shape[0]
|
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):
|
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):
|
def _float(self):
|
||||||
@@ -126,9 +161,17 @@ def _requires_grad(self, requires_grad=True):
|
|||||||
|
|
||||||
|
|
||||||
def _backward(self, *args, **kwargs):
|
def _backward(self, *args, **kwargs):
|
||||||
raise RuntimeError(
|
if args or kwargs:
|
||||||
"loss.backward() is not supported by the MLX backend; use torchmlx.Trainer"
|
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):
|
def _mask_indices(mask):
|
||||||
@@ -144,9 +187,17 @@ def _torch_getitem(self, key):
|
|||||||
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
||||||
indices = _mask_indices(key)
|
indices = _mask_indices(key)
|
||||||
if key.shape == self.shape:
|
if key.shape == self.shape:
|
||||||
return _getitem(self.reshape(-1), indices)
|
result = _getitem(self.reshape(-1), indices)
|
||||||
return _getitem(self, indices)
|
operation = lambda value: _getitem(_reshape(value, (-1,)), indices)
|
||||||
return _getitem(self, key)
|
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):
|
def _torch_setitem(self, key, value):
|
||||||
@@ -163,8 +214,11 @@ def _torch_setitem(self, key, value):
|
|||||||
|
|
||||||
def install(device_type):
|
def install(device_type):
|
||||||
mx.array.transpose = _torch_transpose
|
mx.array.transpose = _torch_transpose
|
||||||
|
mx.array.reshape = _torch_reshape
|
||||||
mx.array.view = _view
|
mx.array.view = _view
|
||||||
mx.array.unsqueeze = _unsqueeze
|
mx.array.unsqueeze = _unsqueeze
|
||||||
|
mx.array.squeeze = _torch_squeeze
|
||||||
|
mx.array.flatten = _flatten
|
||||||
mx.array.float = _float
|
mx.array.float = _float
|
||||||
mx.array.bool = _bool
|
mx.array.bool = _bool
|
||||||
mx.array.pow = _pow
|
mx.array.pow = _pow
|
||||||
@@ -180,6 +234,7 @@ def install(device_type):
|
|||||||
mx.array.contiguous = _contiguous
|
mx.array.contiguous = _contiguous
|
||||||
mx.array.requires_grad_ = _requires_grad
|
mx.array.requires_grad_ = _requires_grad
|
||||||
mx.array.backward = _backward
|
mx.array.backward = _backward
|
||||||
|
mx.array.item = _torch_item
|
||||||
mx.array.device = property(lambda self: device_type("mps"))
|
mx.array.device = property(lambda self: device_type("mps"))
|
||||||
mx.array.__getitem__ = _torch_getitem
|
mx.array.__getitem__ = _torch_getitem
|
||||||
mx.array.__setitem__ = _torch_setitem
|
mx.array.__setitem__ = _torch_setitem
|
||||||
@@ -23,15 +23,32 @@ else:
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as _native
|
import mlx.nn as _native
|
||||||
|
|
||||||
|
class _ParameterTree(dict):
|
||||||
|
def __init__(self, values, model):
|
||||||
|
super().__init__(values)
|
||||||
|
self.model = model
|
||||||
|
|
||||||
class Module(_native.Module):
|
class Module(_native.Module):
|
||||||
def __call__(self, *args, **kwargs):
|
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):
|
def forward(self, *args, **kwargs):
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
f"Module [{type(self).__name__}] is missing the required forward function"
|
f"Module [{type(self).__name__}] is missing the required forward function"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def parameters(self):
|
||||||
|
return _ParameterTree(super().parameters(), self)
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
def to(self, *args, **kwargs):
|
||||||
dtype = kwargs.pop("dtype", None)
|
dtype = kwargs.pop("dtype", None)
|
||||||
device = kwargs.pop("device", None)
|
device = kwargs.pop("device", None)
|
||||||
|
|||||||
@@ -27,6 +27,25 @@ else:
|
|||||||
reduction="mean",
|
reduction="mean",
|
||||||
label_smoothing=0.0,
|
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:
|
if size_average is not None or reduce is not None:
|
||||||
unsupported("torchmlx.nn.functional.cross_entropy legacy reductions")
|
unsupported("torchmlx.nn.functional.cross_entropy legacy reductions")
|
||||||
if input.ndim < 2:
|
if input.ndim < 2:
|
||||||
@@ -56,15 +75,15 @@ else:
|
|||||||
losses = (1 - label_smoothing) * target_losses + label_smoothing * smooth_losses
|
losses = (1 - label_smoothing) * target_losses + label_smoothing * smooth_losses
|
||||||
losses = mx.where(valid, losses, mx.zeros_like(losses))
|
losses = mx.where(valid, losses, mx.zeros_like(losses))
|
||||||
if reduction == "none":
|
if reduction == "none":
|
||||||
return losses.reshape(output_shape)
|
return finish(losses.reshape(output_shape))
|
||||||
if reduction == "sum":
|
if reduction == "sum":
|
||||||
return mx.sum(losses)
|
return finish(mx.sum(losses))
|
||||||
if reduction == "mean":
|
if reduction == "mean":
|
||||||
if weight is None:
|
if weight is None:
|
||||||
denominator = mx.sum(valid)
|
denominator = mx.sum(valid)
|
||||||
else:
|
else:
|
||||||
denominator = mx.sum(mx.where(valid, weight[safe_targets], 0))
|
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}")
|
raise ValueError(f"invalid reduction {reduction!r}")
|
||||||
|
|
||||||
def scaled_dot_product_attention(
|
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:
|
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")
|
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._parameters = params
|
||||||
|
self._model = model
|
||||||
self._optimizer = _native.AdamW(
|
self._optimizer = _native.AdamW(
|
||||||
learning_rate=lr,
|
learning_rate=lr,
|
||||||
betas=list(betas),
|
betas=list(betas),
|
||||||
@@ -37,16 +43,27 @@ else:
|
|||||||
weight_decay=weight_decay,
|
weight_decay=weight_decay,
|
||||||
bias_correction=True,
|
bias_correction=True,
|
||||||
)
|
)
|
||||||
|
self._pending_update = None
|
||||||
|
self._random_before_forward = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def state(self):
|
def state(self):
|
||||||
return self._optimizer.state
|
return self._optimizer.state
|
||||||
|
|
||||||
def zero_grad(self, *args, **kwargs):
|
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):
|
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):
|
def __getattr__(name):
|
||||||
unsupported(f"torchmlx.optim.{name}")
|
unsupported(f"torchmlx.optim.{name}")
|
||||||
|
|||||||
Reference in new issue
Block a user