mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
clean up torchmlx internals
This commit is contained in:
1 parent
4eb86111a8
commit
cd140fae7a
19 files changed
+1115
-794
No files matched your search
@@ -0,0 +1,10 @@
|
||||
# Project instructions
|
||||
|
||||
- Keep code minimal, reliable, comment-free, and focused on user experience.
|
||||
- Do not add tests. Validate with `examples/tinystories-llm/train.py` using the existing uv environment.
|
||||
- Keep the same user code working comfortably on macOS with MLX and Colab with PyTorch.
|
||||
- Treat TorchMLX as an educational experiment and state that clearly in the short README.
|
||||
- Do not use em dashes anywhere in the project.
|
||||
- Work on `main`. Do not commit or push unless explicitly asked.
|
||||
- Use short lowercase commit messages.
|
||||
- Expose only PyTorch-compatible user concepts. Do not introduce TorchMLX-specific abstractions, workflows, or configuration objects that users must learn. Backend machinery must remain private.
|
||||
@@ -12,3 +12,8 @@ from torchmlx import nn, optim
|
||||
```
|
||||
|
||||
see the [tinystories example](examples/tinystories-llm/train.py) and [compatibility details](docs/compatibility.md).
|
||||
|
||||
| backend | 5,000 steps + generation | effective steps/s |
|
||||
| --- | ---: | ---: |
|
||||
| torchmlx mlx, m3 pro | 170.27 s | 29.4 |
|
||||
| pytorch mps, m3 pro | 184.42 s | 27.1 |
|
||||
@@ -4,7 +4,7 @@ TorchMLX targets the common transformer operations used by GPT-2, Llama 3, Qwen
|
||||
|
||||
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.
|
||||
On MLX, TorchMLX tensors are an internal MLX array subclass. Native MLX arrays and their methods are not modified. Model parameters remain native MLX arrays internally and are presented as TorchMLX tensors through the PyTorch-shaped interface.
|
||||
|
||||
Boolean expert routing and `unique` execute eagerly because their output shapes control Python flow.
|
||||
|
||||
@@ -18,8 +18,12 @@ loss.backward()
|
||||
optimizer.step()
|
||||
```
|
||||
|
||||
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.
|
||||
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()` after `backward` and before `step` materializes gradients through a separate compiled path. Models with eager data-dependent operations fall back to an uncompiled replay.
|
||||
|
||||
Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation.
|
||||
The supported subset includes tensor construction, dtype conversion, indexing, arithmetic, comparisons, reshape and view operations, transpose, squeeze and unsqueeze, flatten, expand, chunk, masking, common transformer math, embeddings, linear layers, layer normalization, dropout, GELU, causal attention, cross entropy, and AdamW with its default feature set.
|
||||
|
||||
Unrecorded loss expressions, losses combining multiple model forwards, gradient hooks, parameter `.grad`, higher-order gradients, serialization, additional optimizers, and additional losses remain unsupported on MLX. Use native PyTorch when a program needs behavior outside this subset.
|
||||
|
||||
TorchMLX selects MLX on Apple silicon and PyTorch elsewhere. `TORCHMLX_BACKEND=torch` and `TORCHMLX_BACKEND=mlx` are internal validation overrides. 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.
|
||||
@@ -20,3 +20,14 @@ build-backend = "hatchling.build"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["src/torchmlx"]
|
||||
|
||||
[tool.pyright]
|
||||
typeCheckingMode = "standard"
|
||||
include = ["src", "examples/tinystories-llm/train.py"]
|
||||
venvPath = "."
|
||||
venv = ".venv"
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"pyright>=1.1.414",
|
||||
]
|
||||
+7
-242
@@ -1,259 +1,25 @@
|
||||
import builtins as _builtins
|
||||
import math as _math
|
||||
import platform as _platform
|
||||
from types import SimpleNamespace as _SimpleNamespace
|
||||
|
||||
from ._backend import BACKEND, unsupported
|
||||
|
||||
|
||||
def current_backend():
|
||||
return BACKEND
|
||||
from ._backend import BACKEND
|
||||
|
||||
|
||||
if BACKEND == "torch":
|
||||
import torch as _native
|
||||
|
||||
Tensor = _native.Tensor
|
||||
tensor = _native.tensor
|
||||
from_numpy = _native.from_numpy
|
||||
arange = _native.arange
|
||||
randint = _native.randint
|
||||
chunk = _native.chunk
|
||||
transpose = _native.transpose
|
||||
float16 = _native.float16
|
||||
float32 = _native.float32
|
||||
float64 = _native.float64
|
||||
bfloat16 = _native.bfloat16
|
||||
int8 = _native.int8
|
||||
int16 = _native.int16
|
||||
int32 = _native.int32
|
||||
int64 = _native.int64
|
||||
uint8 = _native.uint8
|
||||
bool = _native.bool
|
||||
half = float16
|
||||
float = float32
|
||||
double = float64
|
||||
int = int32
|
||||
long = int64
|
||||
|
||||
def categorical(logits, dim=-1, num_samples=1):
|
||||
probabilities = _native.softmax(logits, dim=dim)
|
||||
return _native.multinomial(probabilities, num_samples=num_samples)
|
||||
|
||||
def __getattr__(name):
|
||||
return getattr(_native, name)
|
||||
|
||||
from ._torch_backend import *
|
||||
from ._torch_backend import fallback as _fallback
|
||||
else:
|
||||
import mlx.core as _native
|
||||
import numpy as _np
|
||||
from ._mlx_backend import *
|
||||
from ._mlx_backend import fallback as _fallback
|
||||
|
||||
Tensor = _native.array
|
||||
float16 = _native.float16
|
||||
float32 = _native.float32
|
||||
float64 = _native.float64
|
||||
bfloat16 = _native.bfloat16
|
||||
int8 = _native.int8
|
||||
int16 = _native.int16
|
||||
int32 = _native.int32
|
||||
int64 = _native.int64
|
||||
uint8 = _native.uint8
|
||||
bool = _native.bool_
|
||||
half = float16
|
||||
float = float32
|
||||
double = float64
|
||||
int = int32
|
||||
long = int64
|
||||
|
||||
pi = _math.pi
|
||||
|
||||
class device:
|
||||
def __init__(self, value):
|
||||
if isinstance(value, device):
|
||||
value = value.type
|
||||
value = str(value)
|
||||
if value != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
self.type = value
|
||||
self.index = None
|
||||
|
||||
def __str__(self):
|
||||
return self.type
|
||||
|
||||
def __repr__(self):
|
||||
return f"device(type={self.type!r})"
|
||||
|
||||
def __eq__(self, other):
|
||||
return str(other) == self.type
|
||||
|
||||
class _MPS:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return _platform.system() == "Darwin" and _platform.machine() == "arm64"
|
||||
|
||||
class _CUDA:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return False
|
||||
|
||||
backends = _SimpleNamespace(mps=_MPS())
|
||||
cuda = _CUDA()
|
||||
|
||||
from ._mlx_tensor import install as _install_tensor_methods
|
||||
|
||||
_install_tensor_methods(device)
|
||||
|
||||
def _check_device(device):
|
||||
if device is not None and str(device) != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
|
||||
def tensor(data, dtype=None, device=None, requires_grad=False, pin_memory=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
if pin_memory:
|
||||
raise RuntimeError("pin_memory is not supported by the MLX backend")
|
||||
if dtype is None and not isinstance(data, (_native.array, _np.ndarray)):
|
||||
kind = _np.asarray(data).dtype.kind
|
||||
if kind in {"i", "u"}:
|
||||
dtype = int64
|
||||
elif kind == "b":
|
||||
dtype = bool
|
||||
return _native.array(data, dtype=dtype)
|
||||
|
||||
def from_numpy(array):
|
||||
if not isinstance(array, _np.ndarray):
|
||||
raise TypeError("from_numpy expects a numpy.ndarray")
|
||||
return _native.array(array)
|
||||
|
||||
def arange(start, end=None, step=1, *, dtype=None, device=None):
|
||||
_check_device(device)
|
||||
if end is None:
|
||||
start, end = 0, start
|
||||
if dtype is None and all(
|
||||
isinstance(value, _builtins.int) for value in (start, end, step)
|
||||
):
|
||||
dtype = int64
|
||||
return _native.arange(start, end, step, dtype=dtype)
|
||||
|
||||
def randint(low, high, size, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
result = _native.random.randint(low, high, shape=size)
|
||||
return result.astype(dtype or int64)
|
||||
|
||||
def chunk(input, chunks, dim=0):
|
||||
return input.chunk(chunks, dim=dim)
|
||||
|
||||
def transpose(input, dim0, dim1):
|
||||
return input.transpose(dim0, dim1)
|
||||
|
||||
def zeros(*size, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
|
||||
return _native.zeros(shape, dtype=dtype or float32)
|
||||
|
||||
def ones(*size, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
|
||||
return _native.ones(shape, dtype=dtype or float32)
|
||||
|
||||
def zeros_like(input, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
return _native.zeros_like(input).astype(dtype or input.dtype)
|
||||
|
||||
def ones_like(input, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
)
|
||||
return _native.ones_like(input).astype(dtype or input.dtype)
|
||||
|
||||
def cat(tensors, dim=0):
|
||||
return _native.concatenate(tensors, axis=dim)
|
||||
|
||||
def softmax(input, dim, dtype=None):
|
||||
value = input.astype(dtype) if dtype is not None else input
|
||||
return _native.softmax(value, axis=dim)
|
||||
|
||||
def triu(input, diagonal=0):
|
||||
return _native.triu(input, k=diagonal)
|
||||
|
||||
def tril(input, diagonal=0):
|
||||
return _native.tril(input, k=diagonal)
|
||||
|
||||
def topk(input, k, dim=None, largest=True, sorted=True):
|
||||
axis = -1 if dim is None else dim
|
||||
indices = _native.argsort(input, axis=axis)
|
||||
indices = _native.flip(indices, axis=axis) if largest else indices
|
||||
slices = [slice(None)] * input.ndim
|
||||
slices[axis] = slice(0, k)
|
||||
indices = indices[tuple(slices)].astype(int64)
|
||||
values = _native.take_along_axis(input, indices, axis=axis)
|
||||
return values, indices
|
||||
|
||||
def unique(input, sorted=True, return_inverse=False, return_counts=False, dim=None):
|
||||
if return_inverse or return_counts or dim is not None:
|
||||
unsupported("torchmlx.unique with non-default options")
|
||||
values = _native.sort(input.reshape(-1))
|
||||
if values.shape[0] < 2:
|
||||
return values
|
||||
keep = _native.concatenate(
|
||||
[_native.array([True]), values[1:] != values[:-1]]
|
||||
)
|
||||
return values[keep]
|
||||
|
||||
exp = _native.exp
|
||||
sin = _native.sin
|
||||
cos = _native.cos
|
||||
tanh = _native.tanh
|
||||
sqrt = _native.sqrt
|
||||
matmul = _native.matmul
|
||||
outer = _native.outer
|
||||
|
||||
def pow(input, exponent):
|
||||
return _native.power(input, exponent)
|
||||
|
||||
def polar(abs, angle):
|
||||
return abs * _native.exp(_native.array(1j) * angle)
|
||||
|
||||
def categorical(logits, dim=-1, num_samples=1):
|
||||
if num_samples == 1:
|
||||
return _native.random.categorical(logits, axis=dim)[..., None].astype(int64)
|
||||
return _native.random.categorical(
|
||||
logits, axis=dim, num_samples=num_samples
|
||||
).astype(int64)
|
||||
|
||||
def __getattr__(name):
|
||||
unsupported(f"torchmlx.{name}")
|
||||
def __getattr__(name):
|
||||
return _fallback(name)
|
||||
|
||||
|
||||
import importlib as _importlib
|
||||
|
||||
nn = _importlib.import_module("torchmlx.nn")
|
||||
optim = _importlib.import_module("torchmlx.optim")
|
||||
from .trainer import Trainer
|
||||
|
||||
__all__ = [
|
||||
"Tensor",
|
||||
"Trainer",
|
||||
"arange",
|
||||
"backends",
|
||||
"bfloat16",
|
||||
@@ -262,7 +28,6 @@ __all__ = [
|
||||
"cat",
|
||||
"chunk",
|
||||
"cuda",
|
||||
"current_backend",
|
||||
"device",
|
||||
"double",
|
||||
"float",
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
from typing import Any, Iterable, Sequence
|
||||
|
||||
from . import nn as nn, optim as optim
|
||||
|
||||
class device:
|
||||
type: str
|
||||
index: int | None
|
||||
def __init__(self, value: str | device) -> None: ...
|
||||
|
||||
class Tensor:
|
||||
shape: tuple[int, ...]
|
||||
ndim: int
|
||||
dtype: Any
|
||||
device: device
|
||||
grad: Any
|
||||
def size(self, dim: int | None = None) -> Any: ...
|
||||
def view(self, *shape: Any) -> Tensor: ...
|
||||
def reshape(self, *shape: Any) -> Tensor: ...
|
||||
def transpose(self, dim0: int, dim1: int) -> Tensor: ...
|
||||
def unsqueeze(self, dim: int) -> Tensor: ...
|
||||
def squeeze(self, dim: int | None = None) -> Tensor: ...
|
||||
def flatten(self, start_dim: int = 0, end_dim: int = -1) -> Tensor: ...
|
||||
def float(self) -> Tensor: ...
|
||||
def bool(self) -> Tensor: ...
|
||||
def pow(self, exponent: Any) -> Tensor: ...
|
||||
def mean(self, dim: Any = None, keepdim: bool = False, dtype: Any = None) -> Tensor: ...
|
||||
def var(self, dim: Any = None, unbiased: bool = True, keepdim: bool = False, *, correction: int | None = None) -> Tensor: ...
|
||||
def chunk(self, chunks: int, dim: int = 0) -> tuple[Tensor, ...]: ...
|
||||
def to(self, *args: Any, dtype: Any = None, device: Any = None, **kwargs: Any) -> Tensor: ...
|
||||
def repeat_interleave(self, repeats: int, dim: int | None = None) -> Tensor: ...
|
||||
def masked_fill(self, mask: Tensor, value: Any) -> Tensor: ...
|
||||
def expand(self, *sizes: Any) -> Tensor: ...
|
||||
def any(self, dim: Any = None, keepdim: bool = False) -> Tensor: ...
|
||||
def contiguous(self, memory_format: Any = None) -> Tensor: ...
|
||||
def requires_grad_(self, requires_grad: bool = True) -> Tensor: ...
|
||||
def backward(self) -> None: ...
|
||||
def register_hook(self, hook: Any) -> Any: ...
|
||||
def item(self) -> Any: ...
|
||||
def __getitem__(self, key: Any) -> Tensor: ...
|
||||
def __setitem__(self, key: Any, value: Any) -> None: ...
|
||||
def __add__(self, other: Any) -> Tensor: ...
|
||||
def __radd__(self, other: Any) -> Tensor: ...
|
||||
def __sub__(self, other: Any) -> Tensor: ...
|
||||
def __rsub__(self, other: Any) -> Tensor: ...
|
||||
def __mul__(self, other: Any) -> Tensor: ...
|
||||
def __rmul__(self, other: Any) -> Tensor: ...
|
||||
def __truediv__(self, other: Any) -> Tensor: ...
|
||||
def __rtruediv__(self, other: Any) -> Tensor: ...
|
||||
def __pow__(self, other: Any) -> Tensor: ...
|
||||
def __matmul__(self, other: Any) -> Tensor: ...
|
||||
def __eq__(self, other: object) -> Tensor: ... # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
def __ne__(self, other: object) -> Tensor: ... # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
def __lt__(self, other: Any) -> Tensor: ...
|
||||
def __le__(self, other: Any) -> Tensor: ...
|
||||
def __gt__(self, other: Any) -> Tensor: ...
|
||||
def __ge__(self, other: Any) -> Tensor: ...
|
||||
|
||||
float16: Any
|
||||
float32: Any
|
||||
float64: Any
|
||||
bfloat16: Any
|
||||
int8: Any
|
||||
int16: Any
|
||||
int32: Any
|
||||
int64: Any
|
||||
uint8: Any
|
||||
half: Any
|
||||
double: Any
|
||||
long: Any
|
||||
pi: float
|
||||
backends: Any
|
||||
cuda: Any
|
||||
|
||||
def tensor(data: Any, dtype: Any = None, device: Any = None, requires_grad: bool = False, pin_memory: bool = False) -> Tensor: ...
|
||||
def from_numpy(array: Any) -> Tensor: ...
|
||||
def arange(start: Any, end: Any = None, step: Any = 1, *, dtype: Any = None, device: Any = None) -> Tensor: ...
|
||||
def randint(low: int, high: int, size: Sequence[int], *, dtype: Any = None, device: Any = None, requires_grad: bool = False) -> Tensor: ...
|
||||
def zeros(*size: Any, dtype: Any = None, device: Any = None, requires_grad: bool = False) -> Tensor: ...
|
||||
def ones(*size: Any, dtype: Any = None, device: Any = None, requires_grad: bool = False) -> Tensor: ...
|
||||
def zeros_like(input: Tensor, *, dtype: Any = None, device: Any = None, requires_grad: bool = False) -> Tensor: ...
|
||||
def ones_like(input: Tensor, *, dtype: Any = None, device: Any = None, requires_grad: bool = False) -> Tensor: ...
|
||||
def cat(tensors: Iterable[Tensor], dim: int = 0) -> Tensor: ...
|
||||
def chunk(input: Tensor, chunks: int, dim: int = 0) -> tuple[Tensor, ...]: ...
|
||||
def transpose(input: Tensor, dim0: int, dim1: int) -> Tensor: ...
|
||||
def softmax(input: Tensor, dim: int, dtype: Any = None) -> Tensor: ...
|
||||
def triu(input: Tensor, diagonal: int = 0) -> Tensor: ...
|
||||
def tril(input: Tensor, diagonal: int = 0) -> Tensor: ...
|
||||
def topk(input: Tensor, k: int, dim: int | None = None, largest: bool = True, sorted: bool = True) -> tuple[Tensor, Tensor]: ...
|
||||
def unique(input: Tensor, sorted: bool = True, return_inverse: bool = False, return_counts: bool = False, dim: int | None = None) -> Tensor: ...
|
||||
def exp(input: Tensor) -> Tensor: ...
|
||||
def sin(input: Tensor) -> Tensor: ...
|
||||
def cos(input: Tensor) -> Tensor: ...
|
||||
def tanh(input: Tensor) -> Tensor: ...
|
||||
def sqrt(input: Tensor) -> Tensor: ...
|
||||
def matmul(input: Tensor, other: Tensor) -> Tensor: ...
|
||||
def outer(input: Tensor, other: Tensor) -> Tensor: ...
|
||||
def pow(input: Tensor, exponent: Any) -> Tensor: ...
|
||||
def polar(abs: Tensor, angle: Tensor) -> Tensor: ...
|
||||
def categorical(logits: Tensor, dim: int = -1, num_samples: int = 1) -> Tensor: ...
|
||||
def __getattr__(name: str) -> Any: ...
|
||||
+232
-249
@@ -1,64 +1,97 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, cast
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
|
||||
_active_optimizer = None
|
||||
_forward_depth = 0
|
||||
_loss_plans = {}
|
||||
_lineage = {}
|
||||
_replayed_losses = {}
|
||||
_suspended = False
|
||||
from ._mlx_tensor import Tensor
|
||||
|
||||
|
||||
_active_optimizer: ContextVar[Any | None] = ContextVar(
|
||||
"torchmlx_active_optimizer", default=None
|
||||
)
|
||||
_forward_depth: ContextVar[int] = ContextVar("torchmlx_forward_depth", default=0)
|
||||
_suspended: ContextVar[bool] = ContextVar("torchmlx_suspended", default=False)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ArraySlot:
|
||||
def __init__(self, index):
|
||||
self.index = index
|
||||
index: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _Lineage:
|
||||
model: Any
|
||||
args: tuple[Any, ...]
|
||||
kwargs: dict[str, Any]
|
||||
transforms: tuple[Any, ...]
|
||||
random_before: tuple[mx.array, ...] | None
|
||||
compile_allowed: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _LossPlan:
|
||||
model: Any
|
||||
cache_key: tuple[Any, ...]
|
||||
dynamic_inputs: tuple[mx.array, ...]
|
||||
template: Any
|
||||
transforms: tuple[Any, ...]
|
||||
rebuild: Callable[..., mx.array]
|
||||
random_before: tuple[mx.array, ...] | None
|
||||
compile_allowed: bool
|
||||
|
||||
|
||||
class _TrainingPlan:
|
||||
def __init__(self, model, optimizer, objective, compile=True):
|
||||
def __init__(
|
||||
self,
|
||||
model: Any,
|
||||
optimizer: Any,
|
||||
objective: Callable[..., mx.array],
|
||||
compile: bool = True,
|
||||
) -> None:
|
||||
self.model = model
|
||||
self.optimizer = optimizer
|
||||
self._value_and_grad = nn.value_and_grad(model, objective)
|
||||
self._compile_enabled = compile
|
||||
self._compiled_backward = None
|
||||
self._compiled_fused = None
|
||||
self._backward_signatures = set()
|
||||
self._fused_signatures = set()
|
||||
self._compiled_backward: Callable[..., Any] | None = None
|
||||
self._compiled_fused: Callable[..., Any] | None = None
|
||||
self._backward_signatures: set[Any] = set()
|
||||
self._fused_signatures: set[Any] = 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._value_and_grad, inputs=backward_state, outputs=backward_state
|
||||
)
|
||||
self._compiled_fused = mx.compile(
|
||||
self._fused,
|
||||
inputs=fused_state,
|
||||
outputs=fused_state,
|
||||
self._fused, inputs=fused_state, outputs=fused_state
|
||||
)
|
||||
|
||||
def _fused(self, *inputs):
|
||||
def _fused(self, *inputs: mx.array) -> mx.array:
|
||||
loss, gradients = self._value_and_grad(*inputs)
|
||||
self.optimizer.update(self.model, gradients)
|
||||
return loss
|
||||
|
||||
def backward(self, *inputs):
|
||||
def backward(self, *inputs: mx.array) -> Any:
|
||||
if self._compile_enabled:
|
||||
assert self._compiled_backward is not None
|
||||
return self._compiled_backward(*inputs)
|
||||
return self._value_and_grad(*inputs)
|
||||
|
||||
def fused(self, *inputs):
|
||||
def fused(self, *inputs: mx.array) -> mx.array:
|
||||
if self._compile_enabled:
|
||||
assert self._compiled_fused is not None
|
||||
return self._compiled_fused(*inputs)
|
||||
return self._fused(*inputs)
|
||||
|
||||
def disable_compile(self):
|
||||
def disable_compile(self) -> None:
|
||||
self._compile_enabled = False
|
||||
|
||||
|
||||
def _partition_arrays(value, arrays):
|
||||
def _partition_arrays(value: Any, arrays: list[mx.array]) -> Any:
|
||||
if isinstance(value, mx.array):
|
||||
slot = _ArraySlot(len(arrays))
|
||||
arrays.append(value)
|
||||
@@ -72,7 +105,7 @@ def _partition_arrays(value, arrays):
|
||||
return value
|
||||
|
||||
|
||||
def _restore_arrays(template, arrays):
|
||||
def _restore_arrays(template: Any, arrays: tuple[mx.array, ...]) -> Any:
|
||||
if isinstance(template, _ArraySlot):
|
||||
return arrays[template.index]
|
||||
if isinstance(template, tuple):
|
||||
@@ -84,7 +117,7 @@ def _restore_arrays(template, arrays):
|
||||
return template
|
||||
|
||||
|
||||
def _template_key(template):
|
||||
def _template_key(template: Any) -> tuple[Any, ...]:
|
||||
if isinstance(template, _ArraySlot):
|
||||
return ("array",)
|
||||
if isinstance(template, tuple):
|
||||
@@ -102,7 +135,7 @@ def _template_key(template):
|
||||
return ("static", type(template).__qualname__, repr(template))
|
||||
|
||||
|
||||
def _copy_tree(value):
|
||||
def _copy_tree(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {key: _copy_tree(item) for key, item in value.items()}
|
||||
if isinstance(value, list):
|
||||
@@ -112,211 +145,198 @@ def _copy_tree(value):
|
||||
return value
|
||||
|
||||
|
||||
def _snapshot_random_state():
|
||||
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
|
||||
def _snapshot_random_state() -> list[mx.array]:
|
||||
random_state = cast(Any, mx.random.state)
|
||||
state = [key + mx.array(0, dtype=key.dtype) for key in random_state]
|
||||
mx.eval(state)
|
||||
return state
|
||||
|
||||
|
||||
def _restore_random_state(state):
|
||||
for current_key, saved_key in zip(mx.random.state, state):
|
||||
def _restore_random_state(state: Any) -> None:
|
||||
for current_key, saved_key in zip(cast(Any, mx.random.state), state):
|
||||
current_key[...] = saved_key
|
||||
mx.eval(mx.random.state)
|
||||
|
||||
|
||||
def _compile_failure(error):
|
||||
def _compile_failure(error: Exception) -> bool:
|
||||
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
|
||||
if (
|
||||
_suspended
|
||||
or _active_optimizer is None
|
||||
or model is not _active_optimizer._model
|
||||
):
|
||||
def begin_forward(model: Any) -> bool:
|
||||
optimizer = _active_optimizer.get()
|
||||
if _suspended.get() or optimizer is None or model is not optimizer._model:
|
||||
return False
|
||||
if _forward_depth:
|
||||
_active_optimizer._recursive_forward = True
|
||||
_forward_depth += 1
|
||||
depth = _forward_depth.get()
|
||||
if depth:
|
||||
optimizer._recursive_forward = True
|
||||
_forward_depth.set(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 note_eager_evaluation() -> None:
|
||||
optimizer = _active_optimizer.get()
|
||||
if _forward_depth.get() and optimizer is not None and not _suspended.get():
|
||||
optimizer._eager_forward = True
|
||||
|
||||
|
||||
def end_forward(model, args, kwargs, output):
|
||||
global _forward_depth
|
||||
_forward_depth -= 1
|
||||
if _forward_depth != 0 or not isinstance(output, mx.array):
|
||||
def end_forward(
|
||||
model: Any, args: tuple[Any, ...], kwargs: dict[str, Any], output: Any
|
||||
) -> None:
|
||||
optimizer = _active_optimizer.get()
|
||||
depth = _forward_depth.get() - 1
|
||||
_forward_depth.set(depth)
|
||||
if depth != 0 or optimizer is None or not isinstance(output, Tensor):
|
||||
return
|
||||
random_before = _active_optimizer._random_before_forward
|
||||
random_before = optimizer._random_before_forward
|
||||
stochastic = any(
|
||||
current is not previous
|
||||
for current, previous in zip(mx.random.state, random_before)
|
||||
for current, previous in zip(cast(Any, mx.random.state), random_before)
|
||||
)
|
||||
descriptor = (
|
||||
output._lineage = _Lineage(
|
||||
model,
|
||||
args,
|
||||
kwargs,
|
||||
(),
|
||||
random_before if stochastic else None,
|
||||
not (
|
||||
_active_optimizer._recursive_forward
|
||||
or _active_optimizer._eager_forward
|
||||
),
|
||||
not (optimizer._recursive_forward or optimizer._eager_forward),
|
||||
)
|
||||
_lineage[id(output)] = (output, descriptor)
|
||||
|
||||
|
||||
def abort_forward():
|
||||
global _forward_depth
|
||||
_forward_depth -= 1
|
||||
def abort_forward() -> None:
|
||||
depth = max(0, _forward_depth.get() - 1)
|
||||
_forward_depth.set(depth)
|
||||
if depth == 0:
|
||||
optimizer = _active_optimizer.get()
|
||||
if optimizer is not None:
|
||||
optimizer._pending_update = None
|
||||
_active_optimizer.set(None)
|
||||
|
||||
|
||||
def activate(optimizer):
|
||||
global _active_optimizer, _forward_depth
|
||||
_active_optimizer = optimizer
|
||||
_forward_depth = 0
|
||||
_loss_plans.clear()
|
||||
_lineage.clear()
|
||||
_replayed_losses.clear()
|
||||
optimizer._random_before_forward = tuple(mx.random.state)
|
||||
def activate(optimizer: Any) -> None:
|
||||
previous = _active_optimizer.get()
|
||||
if previous is not None and previous is not optimizer:
|
||||
previous._pending_update = None
|
||||
_active_optimizer.set(optimizer)
|
||||
_forward_depth.set(0)
|
||||
optimizer._random_before_forward = tuple(cast(Any, mx.random.state))
|
||||
optimizer._recursive_forward = False
|
||||
optimizer._eager_forward = False
|
||||
|
||||
|
||||
def propagate(source, result, operation, signature, operands=()):
|
||||
if _suspended:
|
||||
def propagate(
|
||||
source: Tensor,
|
||||
result: Tensor,
|
||||
operation: Callable[..., mx.array],
|
||||
signature: tuple[Any, ...],
|
||||
operands: tuple[mx.array, ...] = (),
|
||||
) -> Tensor:
|
||||
if _suspended.get() or source._lineage is None:
|
||||
return result
|
||||
entry = _lineage.get(id(source))
|
||||
if entry is None or entry[0] is not source:
|
||||
return result
|
||||
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
|
||||
_lineage[id(result)] = (
|
||||
result,
|
||||
(
|
||||
model,
|
||||
args,
|
||||
kwargs,
|
||||
transforms + ((operation, signature, operands),),
|
||||
random_before,
|
||||
compile_allowed,
|
||||
),
|
||||
lineage: _Lineage = source._lineage
|
||||
result._lineage = _Lineage(
|
||||
lineage.model,
|
||||
lineage.args,
|
||||
lineage.kwargs,
|
||||
lineage.transforms + ((operation, signature, operands),),
|
||||
lineage.random_before,
|
||||
lineage.compile_allowed,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _plan_for(
|
||||
optimizer, model, cache_key, template, transforms, rebuild, compile_allowed
|
||||
):
|
||||
plan = optimizer._training_plans.get(cache_key)
|
||||
def _plan_for(optimizer: Any, loss_plan: _LossPlan) -> _TrainingPlan:
|
||||
plan = optimizer._training_plans.get(loss_plan.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)
|
||||
def objective(*current_inputs: mx.array) -> mx.array:
|
||||
current_args, current_kwargs, transform_operands, current_operands = (
|
||||
_restore_arrays(loss_plan.template, current_inputs)
|
||||
)
|
||||
output = model(*current_args, **current_kwargs)
|
||||
for (operation, _, _), operation_operands in zip(
|
||||
transforms, current_transform_operands
|
||||
output = loss_plan.model(*current_args, **current_kwargs)
|
||||
for (operation, _, _), operands in zip(
|
||||
loss_plan.transforms, transform_operands
|
||||
):
|
||||
output = operation(output, *operation_operands)
|
||||
return rebuild(output, *current_operands)
|
||||
output = operation(output, *operands)
|
||||
return loss_plan.rebuild(output, *current_operands)
|
||||
|
||||
plan = _TrainingPlan(
|
||||
model, optimizer._optimizer, objective, compile=compile_allowed
|
||||
loss_plan.model,
|
||||
optimizer._optimizer,
|
||||
objective,
|
||||
compile=loss_plan.compile_allowed,
|
||||
)
|
||||
optimizer._training_plans[cache_key] = plan
|
||||
optimizer._training_plans[loss_plan.cache_key] = plan
|
||||
return plan
|
||||
|
||||
|
||||
def register_loss(loss, loss_input, rebuild, operands, signature):
|
||||
if _suspended or _active_optimizer is None:
|
||||
def register_loss(
|
||||
loss: Tensor,
|
||||
loss_input: Tensor,
|
||||
rebuild: Callable[..., mx.array],
|
||||
operands: tuple[mx.array, ...],
|
||||
signature: tuple[Any, ...],
|
||||
) -> Tensor:
|
||||
optimizer = _active_optimizer.get()
|
||||
if _suspended.get() or optimizer is None or loss_input._lineage 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, transforms, random_before, compile_allowed = entry[1]
|
||||
lineage: _Lineage = loss_input._lineage
|
||||
context = (
|
||||
args,
|
||||
kwargs,
|
||||
tuple(transform_operands for _, _, transform_operands in transforms),
|
||||
lineage.args,
|
||||
lineage.kwargs,
|
||||
tuple(cast(Any, item)[2] for item in lineage.transforms),
|
||||
operands,
|
||||
)
|
||||
dynamic_inputs = []
|
||||
dynamic_inputs: list[mx.array] = []
|
||||
template = _partition_arrays(context, dynamic_inputs)
|
||||
cache_key = (
|
||||
getattr(model, "training", None),
|
||||
getattr(lineage.model, "training", None),
|
||||
_template_key(template),
|
||||
tuple(transform_signature for _, transform_signature, _ in transforms),
|
||||
tuple(item[1] for item in lineage.transforms),
|
||||
signature,
|
||||
compile_allowed,
|
||||
lineage.compile_allowed,
|
||||
)
|
||||
_loss_plans[id(loss)] = (
|
||||
loss,
|
||||
model,
|
||||
loss._loss_plan = _LossPlan(
|
||||
lineage.model,
|
||||
cache_key,
|
||||
tuple(dynamic_inputs),
|
||||
template,
|
||||
transforms,
|
||||
lineage.transforms,
|
||||
rebuild,
|
||||
random_before,
|
||||
compile_allowed,
|
||||
lineage.random_before,
|
||||
lineage.compile_allowed,
|
||||
)
|
||||
return loss
|
||||
|
||||
|
||||
def backward(loss):
|
||||
if _active_optimizer is None:
|
||||
def backward(loss: Tensor) -> None:
|
||||
optimizer = _active_optimizer.get()
|
||||
if 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:
|
||||
loss_plan = loss._loss_plan
|
||||
if loss_plan is None:
|
||||
_clear(optimizer)
|
||||
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 one supported loss from one model 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 = (
|
||||
try:
|
||||
plan = _plan_for(optimizer, loss_plan)
|
||||
except Exception:
|
||||
_clear(optimizer)
|
||||
raise
|
||||
optimizer._pending_update = (
|
||||
"staged",
|
||||
loss,
|
||||
plan,
|
||||
dynamic_inputs,
|
||||
random_before,
|
||||
loss_plan.dynamic_inputs,
|
||||
loss_plan.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):
|
||||
def _restore_tree(current: Any, saved: Any) -> Any:
|
||||
if isinstance(current, dict) and isinstance(saved, dict):
|
||||
for key in tuple(current):
|
||||
if key not in saved:
|
||||
@@ -330,16 +350,26 @@ def _restore_tree(current, saved):
|
||||
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)
|
||||
]
|
||||
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
|
||||
def _restore_execution(
|
||||
model: Any, optimizer: Any, parameters: Any, optimizer_state: Any
|
||||
) -> None:
|
||||
model.update(parameters)
|
||||
if optimizer_state is not None:
|
||||
_restore_tree(optimizer.state, optimizer_state)
|
||||
|
||||
|
||||
def _execute(
|
||||
plan: _TrainingPlan,
|
||||
inputs: tuple[mx.array, ...],
|
||||
fused: bool,
|
||||
random_before: Any,
|
||||
) -> Any:
|
||||
model = plan.model
|
||||
optimizer = plan.optimizer
|
||||
signatures = plan._fused_signatures if fused else plan._backward_signatures
|
||||
@@ -358,22 +388,19 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
|
||||
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
|
||||
token = _suspended.set(True)
|
||||
try:
|
||||
try:
|
||||
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
|
||||
current_parameters = model.parameters()
|
||||
current_parameters = model._native_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 [],
|
||||
mx.random.state if random_before is not None else [],
|
||||
)
|
||||
signatures.add(signature)
|
||||
if fused:
|
||||
@@ -385,19 +412,15 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
|
||||
_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_parameters = model._native_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 [],
|
||||
mx.random.state if random_before is not None else [],
|
||||
)
|
||||
if fused:
|
||||
optimizer._torchmlx_parameters = current_parameters
|
||||
@@ -405,105 +428,65 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
|
||||
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
|
||||
finally:
|
||||
_suspended = False
|
||||
_suspended.reset(token)
|
||||
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
|
||||
def _materialize(optimizer: Any) -> mx.array:
|
||||
_, loss, plan, dynamic_inputs, random_before = optimizer._pending_update
|
||||
replayed_loss, gradients = _execute(
|
||||
plan, dynamic_inputs, fused=False, random_before=random_before
|
||||
)
|
||||
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
||||
loss._replayed_value = replayed_loss
|
||||
optimizer._pending_update = ("materialized", loss, plan.model, gradients)
|
||||
return replayed_loss
|
||||
|
||||
|
||||
def replayed_value(value):
|
||||
if _active_optimizer is not None:
|
||||
pending = _active_optimizer._pending_update
|
||||
def replayed_value(value: Tensor) -> mx.array:
|
||||
optimizer = _active_optimizer.get()
|
||||
if optimizer is not None:
|
||||
pending = 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))
|
||||
if entry is None or entry[0] is not value:
|
||||
return value
|
||||
return entry[1]
|
||||
return _materialize(optimizer)
|
||||
return value._replayed_value if value._replayed_value is not None else value
|
||||
|
||||
|
||||
def _finish_step(loss, replayed_loss):
|
||||
global _active_optimizer, _forward_depth
|
||||
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
||||
_active_optimizer._pending_update = None
|
||||
_active_optimizer = None
|
||||
_forward_depth = 0
|
||||
_loss_plans.clear()
|
||||
_lineage.clear()
|
||||
def _clear(optimizer: Any) -> None:
|
||||
optimizer._pending_update = None
|
||||
_active_optimizer.set(None)
|
||||
_forward_depth.set(0)
|
||||
|
||||
|
||||
def step(optimizer):
|
||||
def step(optimizer: Any) -> None:
|
||||
pending = optimizer._pending_update
|
||||
if pending is None:
|
||||
_clear(optimizer)
|
||||
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
|
||||
try:
|
||||
if pending[0] == "staged":
|
||||
_, loss, plan, dynamic_inputs, random_before = pending
|
||||
replayed_loss = _execute(
|
||||
plan, dynamic_inputs, fused=True, random_before=random_before
|
||||
)
|
||||
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,
|
||||
)
|
||||
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._native_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 = loss._replayed_value
|
||||
loss._replayed_value = replayed_loss
|
||||
finally:
|
||||
_clear(optimizer)
|
||||
@@ -0,0 +1,219 @@
|
||||
import builtins
|
||||
import math
|
||||
import platform
|
||||
from types import SimpleNamespace
|
||||
|
||||
import mlx.core as mx
|
||||
import numpy as np
|
||||
|
||||
from ._backend import unsupported
|
||||
from ._mlx_tensor import Tensor, wrap
|
||||
|
||||
|
||||
float16 = mx.float16
|
||||
float32 = mx.float32
|
||||
float64 = mx.float64
|
||||
bfloat16 = mx.bfloat16
|
||||
int8 = mx.int8
|
||||
int16 = mx.int16
|
||||
int32 = mx.int32
|
||||
int64 = mx.int64
|
||||
uint8 = mx.uint8
|
||||
bool = mx.bool_
|
||||
half = float16
|
||||
float = float32
|
||||
double = float64
|
||||
int = int32
|
||||
long = int64
|
||||
pi = math.pi
|
||||
|
||||
|
||||
class device:
|
||||
def __init__(self, value):
|
||||
if isinstance(value, device):
|
||||
value = value.type
|
||||
value = str(value)
|
||||
if value != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
self.type = value
|
||||
self.index = None
|
||||
|
||||
def __str__(self):
|
||||
return self.type
|
||||
|
||||
def __repr__(self):
|
||||
return f"device(type={self.type!r})"
|
||||
|
||||
def __eq__(self, other):
|
||||
return str(other) == self.type
|
||||
|
||||
|
||||
class _MPS:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return platform.system() == "Darwin" and platform.machine() == "arm64"
|
||||
|
||||
|
||||
class _CUDA:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return False
|
||||
|
||||
|
||||
backends = SimpleNamespace(mps=_MPS())
|
||||
cuda = _CUDA()
|
||||
|
||||
|
||||
def _check_device(value):
|
||||
if value is not None and str(value) != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
|
||||
|
||||
def tensor(data, dtype=None, device=None, requires_grad=False, pin_memory=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
if pin_memory:
|
||||
raise RuntimeError("pin_memory is not supported by the MLX backend")
|
||||
if dtype is None and not isinstance(data, (mx.array, np.ndarray)):
|
||||
kind = np.asarray(data).dtype.kind
|
||||
if kind in {"i", "u"}:
|
||||
dtype = int64
|
||||
elif kind == "b":
|
||||
dtype = bool
|
||||
return wrap(mx.array(data, dtype=dtype))
|
||||
|
||||
|
||||
def from_numpy(array):
|
||||
if not isinstance(array, np.ndarray):
|
||||
raise TypeError("from_numpy expects a numpy.ndarray")
|
||||
return wrap(mx.array(array))
|
||||
|
||||
|
||||
def arange(start, end=None, step=1, *, dtype=None, device=None):
|
||||
_check_device(device)
|
||||
if end is None:
|
||||
start, end = 0, start
|
||||
if dtype is None and all(isinstance(value, builtins.int) for value in (start, end, step)):
|
||||
dtype = int64
|
||||
return wrap(mx.arange(start, end, step, dtype=dtype))
|
||||
|
||||
|
||||
def randint(low, high, size, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
return wrap(mx.random.randint(low, high, shape=size).astype(dtype or int64))
|
||||
|
||||
|
||||
def chunk(input, chunks, dim=0):
|
||||
return input.chunk(chunks, dim=dim)
|
||||
|
||||
|
||||
def transpose(input, dim0, dim1):
|
||||
return input.transpose(dim0, dim1)
|
||||
|
||||
|
||||
def zeros(*size, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
|
||||
return wrap(mx.zeros(shape, dtype=dtype or float32))
|
||||
|
||||
|
||||
def ones(*size, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
|
||||
return wrap(mx.ones(shape, dtype=dtype or float32))
|
||||
|
||||
|
||||
def zeros_like(input, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
return wrap(mx.zeros_like(input).astype(dtype or input.dtype))
|
||||
|
||||
|
||||
def ones_like(input, *, dtype=None, device=None, requires_grad=False):
|
||||
_check_device(device)
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad is not supported by the MLX backend")
|
||||
return wrap(mx.ones_like(input).astype(dtype or input.dtype))
|
||||
|
||||
|
||||
def cat(tensors, dim=0):
|
||||
return wrap(mx.concatenate(tensors, axis=dim))
|
||||
|
||||
|
||||
def softmax(input, dim, dtype=None):
|
||||
value = input.astype(dtype) if dtype is not None else input
|
||||
return wrap(mx.softmax(value, axis=dim))
|
||||
|
||||
|
||||
def triu(input, diagonal=0):
|
||||
return wrap(mx.triu(input, k=diagonal))
|
||||
|
||||
|
||||
def tril(input, diagonal=0):
|
||||
return wrap(mx.tril(input, k=diagonal))
|
||||
|
||||
|
||||
def topk(input, k, dim=None, largest=True, sorted=True):
|
||||
axis = -1 if dim is None else dim
|
||||
indices = mx.argsort(input, axis=axis)
|
||||
indices = mx.flip(indices, axis=axis) if largest else indices
|
||||
slices = [slice(None)] * input.ndim
|
||||
slices[axis] = slice(0, k)
|
||||
indices = indices[tuple(slices)].astype(int64)
|
||||
values = mx.take_along_axis(input, indices, axis=axis)
|
||||
return wrap((values, indices))
|
||||
|
||||
|
||||
def unique(input, sorted=True, return_inverse=False, return_counts=False, dim=None):
|
||||
if return_inverse or return_counts or dim is not None:
|
||||
unsupported("torchmlx.unique with non-default options")
|
||||
values = mx.sort(input.reshape(-1))
|
||||
if values.shape[0] < 2:
|
||||
return wrap(values)
|
||||
keep = mx.concatenate([mx.array([True]), values[1:] != values[:-1]])
|
||||
return wrap(values[keep])
|
||||
|
||||
|
||||
def _unary(function):
|
||||
return lambda input: wrap(function(input))
|
||||
|
||||
|
||||
exp = _unary(mx.exp)
|
||||
sin = _unary(mx.sin)
|
||||
cos = _unary(mx.cos)
|
||||
tanh = _unary(mx.tanh)
|
||||
sqrt = _unary(mx.sqrt)
|
||||
|
||||
|
||||
def matmul(input, other):
|
||||
return wrap(mx.matmul(input, other))
|
||||
|
||||
|
||||
def outer(input, other):
|
||||
return wrap(mx.outer(input, other))
|
||||
|
||||
|
||||
def pow(input, exponent):
|
||||
return wrap(mx.power(input, exponent))
|
||||
|
||||
|
||||
def polar(abs, angle):
|
||||
return wrap(abs * mx.exp(mx.array(1j) * angle))
|
||||
|
||||
|
||||
def categorical(logits, dim=-1, num_samples=1):
|
||||
if num_samples == 1:
|
||||
return wrap(mx.random.categorical(logits, axis=dim)[..., None].astype(int64))
|
||||
return wrap(mx.random.categorical(logits, axis=dim, num_samples=num_samples).astype(int64))
|
||||
|
||||
|
||||
def fallback(name):
|
||||
unsupported(f"torchmlx.{name}")
|
||||
+311
-253
@@ -1,275 +1,333 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable, cast
|
||||
|
||||
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
|
||||
_native_getitem = mx.array.__getitem__
|
||||
_native_setitem = mx.array.__setitem__
|
||||
_native_item = mx.array.item
|
||||
_native_reshape = mx.array.reshape
|
||||
_native_squeeze = mx.array.squeeze
|
||||
_native_transpose = mx.array.transpose
|
||||
|
||||
|
||||
def _torch_transpose(self, dim0=None, dim1=None):
|
||||
if dim1 is None:
|
||||
if dim0 is None:
|
||||
axes = list(reversed(range(self.ndim)))
|
||||
def wrap(value: Any) -> Any:
|
||||
if isinstance(value, Tensor):
|
||||
return value
|
||||
if isinstance(value, mx.array):
|
||||
return Tensor(value)
|
||||
if isinstance(value, tuple):
|
||||
return tuple(wrap(item) for item in value)
|
||||
if isinstance(value, list):
|
||||
return [wrap(item) for item in value]
|
||||
if isinstance(value, dict):
|
||||
return {key: wrap(item) for key, item in value.items()}
|
||||
return value
|
||||
|
||||
|
||||
class Tensor(mx.array):
|
||||
_lineage: Any = None
|
||||
_loss_plan: Any = None
|
||||
_replayed_value: mx.array | None = None
|
||||
|
||||
def _result(
|
||||
self,
|
||||
value: Any,
|
||||
operation: Callable[..., mx.array] | None = None,
|
||||
signature: tuple[Any, ...] = (),
|
||||
operands: tuple[mx.array, ...] = (),
|
||||
) -> Any:
|
||||
result = wrap(value)
|
||||
if operation is not None and isinstance(result, Tensor):
|
||||
from ._autograd import propagate
|
||||
|
||||
propagate(self, result, operation, signature, operands)
|
||||
return result
|
||||
|
||||
@property
|
||||
def device(self) -> Any:
|
||||
from ._mlx_backend import device
|
||||
|
||||
return device("mps")
|
||||
|
||||
@property
|
||||
def grad(self) -> Any:
|
||||
raise RuntimeError("parameter gradients are not exposed by the MLX backend")
|
||||
|
||||
def register_hook(self, hook: Any) -> Any:
|
||||
raise RuntimeError("gradient hooks are not supported by the MLX backend")
|
||||
|
||||
def transpose(self, dim0: Any = None, dim1: int | None = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
if dim1 is None:
|
||||
axes = list(reversed(range(self.ndim))) if dim0 is None else dim0
|
||||
else:
|
||||
axes = dim0
|
||||
else:
|
||||
axes = list(range(self.ndim))
|
||||
axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
|
||||
result = _transpose(self, axes)
|
||||
from ._autograd import propagate
|
||||
|
||||
return propagate(
|
||||
self,
|
||||
result,
|
||||
lambda value: _transpose(value, axes),
|
||||
("transpose", tuple(axes)),
|
||||
)
|
||||
|
||||
|
||||
def _view(self, *shape):
|
||||
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
||||
shape = shape[0]
|
||||
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),
|
||||
("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),
|
||||
("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),
|
||||
("squeeze", 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):
|
||||
return self.astype(mx.float32)
|
||||
|
||||
|
||||
def _bool(self):
|
||||
return self.astype(mx.bool_)
|
||||
|
||||
|
||||
def _pow(self, exponent):
|
||||
return mx.power(self, exponent)
|
||||
|
||||
|
||||
def _torch_mean(self, dim=None, keepdim=False, dtype=None):
|
||||
value = self.astype(dtype) if dtype is not None else self
|
||||
return _mean(value, axis=dim, keepdims=keepdim)
|
||||
|
||||
|
||||
def _torch_var(
|
||||
self,
|
||||
dim=None,
|
||||
unbiased=True,
|
||||
keepdim=False,
|
||||
*,
|
||||
correction=None,
|
||||
):
|
||||
ddof = int(unbiased) if correction is None else correction
|
||||
return _var(self, axis=dim, keepdims=keepdim, ddof=ddof)
|
||||
|
||||
|
||||
def _size(self, dim=None):
|
||||
return self.shape if dim is None else self.shape[dim]
|
||||
|
||||
|
||||
def _chunk(self, chunks, dim=0):
|
||||
if chunks <= 0:
|
||||
raise ValueError("chunks must be greater than 0")
|
||||
length = self.shape[dim]
|
||||
if length == 0:
|
||||
return tuple(mx.split(self, chunks, axis=dim))
|
||||
chunk_size = (length + chunks - 1) // chunks
|
||||
indices = list(range(chunk_size, length, chunk_size))
|
||||
return tuple(mx.split(self, indices, axis=dim))
|
||||
|
||||
|
||||
def _to(self, *args, dtype=None, device=None, **kwargs):
|
||||
if kwargs:
|
||||
name = next(iter(kwargs))
|
||||
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
|
||||
for value in args:
|
||||
if isinstance(value, mx.Dtype):
|
||||
dtype = value
|
||||
elif isinstance(value, mx.array):
|
||||
dtype = value.dtype
|
||||
else:
|
||||
device = value
|
||||
if device is not None and str(device) != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
return self.astype(dtype) if dtype is not None and dtype != self.dtype else self
|
||||
|
||||
|
||||
def _repeat_interleave(self, repeats, dim=None):
|
||||
return mx.repeat(self, repeats, axis=dim)
|
||||
|
||||
|
||||
def _masked_fill(self, mask, value):
|
||||
return mx.where(mask, mx.array(value, dtype=self.dtype), self)
|
||||
|
||||
|
||||
def _expand(self, *sizes):
|
||||
if len(sizes) == 1 and isinstance(sizes[0], (tuple, list)):
|
||||
sizes = tuple(sizes[0])
|
||||
if len(sizes) < self.ndim:
|
||||
raise ValueError("expanded size must have at least as many dimensions as the tensor")
|
||||
source = (1,) * (len(sizes) - self.ndim) + self.shape
|
||||
target = tuple(current if requested == -1 else requested for requested, current in zip(sizes, source))
|
||||
return mx.broadcast_to(self.reshape(source), target)
|
||||
|
||||
|
||||
def _torch_any(self, dim=None, keepdim=False):
|
||||
return _any(self, axis=dim, keepdims=keepdim)
|
||||
|
||||
|
||||
def _contiguous(self, memory_format=None):
|
||||
return self
|
||||
|
||||
|
||||
def _requires_grad(self, requires_grad=True):
|
||||
if requires_grad:
|
||||
raise RuntimeError(
|
||||
"requires_grad_ is not supported by the MLX backend; use torchmlx.Trainer"
|
||||
axes = list(range(self.ndim))
|
||||
axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
|
||||
return self._result(
|
||||
_native_transpose(self, axes),
|
||||
lambda value: _native_transpose(value, axes),
|
||||
("transpose", tuple(axes)),
|
||||
)
|
||||
return self
|
||||
|
||||
def view(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
return self.reshape(*shape)
|
||||
|
||||
def _backward(self, *args, **kwargs):
|
||||
if args or kwargs:
|
||||
raise TypeError("MLX backward compatibility does not accept arguments")
|
||||
from ._autograd import backward
|
||||
def reshape(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
||||
shape = tuple(shape[0])
|
||||
return self._result(
|
||||
_native_reshape(self, shape),
|
||||
lambda value: _native_reshape(value, shape),
|
||||
("reshape", tuple(shape)),
|
||||
)
|
||||
|
||||
backward(self)
|
||||
def unsqueeze(self, dim: int) -> Tensor:
|
||||
return self._result(
|
||||
mx.expand_dims(self, axis=dim),
|
||||
lambda value: mx.expand_dims(value, axis=dim),
|
||||
("unsqueeze", dim),
|
||||
)
|
||||
|
||||
def squeeze(self, dim: int | None = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
return self._result(
|
||||
_native_squeeze(self, axis=dim),
|
||||
lambda value: _native_squeeze(value, axis=dim),
|
||||
("squeeze", dim),
|
||||
)
|
||||
|
||||
def _torch_item(self):
|
||||
from ._autograd import note_eager_evaluation, replayed_value
|
||||
def flatten(self, start_dim: int = 0, end_dim: int = -1) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
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 self.reshape(shape)
|
||||
|
||||
note_eager_evaluation()
|
||||
return _item(replayed_value(self))
|
||||
def float(self) -> Tensor:
|
||||
return self.astype(mx.float32)
|
||||
|
||||
def bool(self) -> Tensor:
|
||||
return self.astype(mx.bool_)
|
||||
|
||||
def _mask_indices(mask):
|
||||
flat = mask.reshape(-1)
|
||||
count = int(mx.sum(flat).item())
|
||||
order = mx.argsort(flat.astype(mx.int32))
|
||||
if count == 0:
|
||||
return order[:0].astype(mx.int64)
|
||||
return order[-count:].astype(mx.int64)
|
||||
def pow(self, exponent: Any) -> Tensor:
|
||||
return self._binary(mx.power, exponent, "pow")
|
||||
|
||||
def mean( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self, dim: Any = None, keepdim: bool = False, dtype: Any = None
|
||||
) -> Tensor:
|
||||
value = mx.astype(self, dtype) if dtype is not None else self
|
||||
return wrap(mx.mean(value, axis=dim, keepdims=keepdim))
|
||||
|
||||
def _torch_getitem(self, key):
|
||||
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
||||
indices = _mask_indices(key)
|
||||
if key.shape == self.shape:
|
||||
result = _getitem(self.reshape(-1), indices)
|
||||
operation = lambda value, current_indices: _getitem(
|
||||
_reshape(value, (-1,)), current_indices
|
||||
def var( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self,
|
||||
dim: Any = None,
|
||||
unbiased: bool = True,
|
||||
keepdim: bool = False,
|
||||
*,
|
||||
correction: int | None = None,
|
||||
) -> Tensor:
|
||||
ddof = int(unbiased) if correction is None else correction
|
||||
return wrap(mx.var(self, axis=dim, keepdims=keepdim, ddof=ddof))
|
||||
|
||||
def size(self, dim: int | None = None) -> Any: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
return self.shape if dim is None else self.shape[dim]
|
||||
|
||||
def chunk(self, chunks: int, dim: int = 0) -> tuple[Tensor, ...]:
|
||||
if chunks <= 0:
|
||||
raise ValueError("chunks must be greater than 0")
|
||||
length = self.shape[dim]
|
||||
if length == 0:
|
||||
return wrap(tuple(mx.split(self, chunks, axis=dim)))
|
||||
chunk_size = (length + chunks - 1) // chunks
|
||||
indices = list(range(chunk_size, length, chunk_size))
|
||||
return wrap(tuple(mx.split(self, indices, axis=dim)))
|
||||
|
||||
def to(
|
||||
self,
|
||||
*args: Any,
|
||||
dtype: Any = None,
|
||||
device: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> Tensor:
|
||||
if kwargs:
|
||||
name = next(iter(kwargs))
|
||||
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
|
||||
for value in args:
|
||||
if isinstance(value, mx.Dtype):
|
||||
dtype = value
|
||||
elif isinstance(value, mx.array):
|
||||
dtype = value.dtype
|
||||
else:
|
||||
device = value
|
||||
if device is not None and str(device) != "mps":
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
return self.astype(dtype) if dtype is not None and dtype != self.dtype else self
|
||||
|
||||
def repeat_interleave(self, repeats: int, dim: int | None = None) -> Tensor:
|
||||
return wrap(mx.repeat(self, repeats, axis=dim))
|
||||
|
||||
def masked_fill(self, mask: mx.array, value: Any) -> Tensor:
|
||||
return wrap(mx.where(mask, mx.array(value, dtype=self.dtype), self))
|
||||
|
||||
def expand(self, *sizes: Any) -> Tensor:
|
||||
if len(sizes) == 1 and isinstance(sizes[0], (tuple, list)):
|
||||
sizes = tuple(sizes[0])
|
||||
if len(sizes) < self.ndim:
|
||||
raise ValueError("expanded size must have at least as many dimensions as the tensor")
|
||||
source = (1,) * (len(sizes) - self.ndim) + self.shape
|
||||
target = tuple(
|
||||
current if requested == -1 else requested
|
||||
for requested, current in zip(sizes, source)
|
||||
)
|
||||
return wrap(mx.broadcast_to(_native_reshape(self, source), target))
|
||||
|
||||
def any(self, dim: Any = None, keepdim: bool = False) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
return wrap(mx.any(self, axis=dim, keepdims=keepdim))
|
||||
|
||||
def contiguous(self, memory_format: Any = None) -> Tensor:
|
||||
return self
|
||||
|
||||
def requires_grad_(self, requires_grad: bool = True) -> Tensor:
|
||||
if requires_grad:
|
||||
raise RuntimeError("requires_grad_ is not supported by the MLX backend")
|
||||
return self
|
||||
|
||||
def backward(self, *args: Any, **kwargs: Any) -> None:
|
||||
if args or kwargs:
|
||||
if kwargs.get("create_graph") or kwargs.get("retain_graph"):
|
||||
raise RuntimeError("higher-order gradients are not supported by the MLX backend")
|
||||
raise TypeError("MLX backward compatibility does not accept arguments")
|
||||
from ._autograd import backward
|
||||
|
||||
backward(self)
|
||||
|
||||
def item(self) -> Any:
|
||||
from ._autograd import note_eager_evaluation, replayed_value
|
||||
|
||||
note_eager_evaluation()
|
||||
return _native_item(replayed_value(self))
|
||||
|
||||
def astype(self, dtype: Any, *, stream: Any = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
kwargs = {} if stream is None else {"stream": stream}
|
||||
return wrap(mx.astype(self, dtype, **kwargs))
|
||||
|
||||
def __getitem__(self, key: Any) -> Tensor:
|
||||
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
||||
indices = _mask_indices(key)
|
||||
if key.shape == self.shape:
|
||||
result = _native_getitem(_native_reshape(self, (-1,)), indices)
|
||||
operation = lambda value, current: _native_getitem(
|
||||
_native_reshape(value, (-1,)), current
|
||||
)
|
||||
signature = ("getitem_bool", "flat")
|
||||
else:
|
||||
result = _native_getitem(self, indices)
|
||||
operation = lambda value, current: _native_getitem(value, current)
|
||||
signature = ("getitem_bool", "first_axis")
|
||||
return self._result(result, operation, signature, (indices,))
|
||||
if isinstance(key, mx.array):
|
||||
return self._result(
|
||||
_native_getitem(self, key),
|
||||
lambda value, current: _native_getitem(value, current),
|
||||
("getitem_array",),
|
||||
(key,),
|
||||
)
|
||||
signature = ("getitem_bool", "flat")
|
||||
else:
|
||||
result = _getitem(self, 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 self._result(
|
||||
_native_getitem(self, key),
|
||||
lambda value: _native_getitem(value, key),
|
||||
("getitem", type(key).__qualname__, repr(key)),
|
||||
)
|
||||
|
||||
return propagate(self, result, operation, signature, operands)
|
||||
|
||||
|
||||
def _torch_setitem(self, key, value):
|
||||
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
||||
indices = _mask_indices(key)
|
||||
if key.shape == self.shape:
|
||||
flat = self.reshape(-1)
|
||||
_setitem(flat, indices, value)
|
||||
def __setitem__(self, key: Any, value: Any) -> None:
|
||||
if isinstance(key, mx.array) and key.dtype == mx.bool_:
|
||||
indices = _mask_indices(key)
|
||||
if key.shape == self.shape:
|
||||
_native_setitem(_native_reshape(self, (-1,)), indices, value)
|
||||
return
|
||||
_native_setitem(self, indices, value)
|
||||
return
|
||||
_setitem(self, indices, value)
|
||||
return
|
||||
_setitem(self, key, value)
|
||||
_native_setitem(self, key, value)
|
||||
|
||||
def _binary(self, function: Callable[..., mx.array], other: Any, name: str) -> Tensor:
|
||||
if isinstance(other, Tensor) and other._lineage is not None:
|
||||
return wrap(function(self, other))
|
||||
operands = (other,) if isinstance(other, mx.array) else ()
|
||||
operation = (
|
||||
(lambda value, current: function(value, current))
|
||||
if operands
|
||||
else (lambda value: function(value, other))
|
||||
)
|
||||
signature = (name, "array" if operands else repr(other))
|
||||
return self._result(function(self, other), operation, signature, operands)
|
||||
|
||||
def _reverse(self, function: Callable[..., mx.array], other: Any) -> Tensor:
|
||||
return wrap(function(other, self))
|
||||
|
||||
def __add__(self, other: Any) -> Tensor:
|
||||
return self._binary(mx.add, other, "add")
|
||||
|
||||
def __radd__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.add, other)
|
||||
|
||||
def __sub__(self, other: Any) -> Tensor:
|
||||
return self._binary(mx.subtract, other, "sub")
|
||||
|
||||
def __rsub__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.subtract, other)
|
||||
|
||||
def __mul__(self, other: Any) -> Tensor:
|
||||
return self._binary(mx.multiply, other, "mul")
|
||||
|
||||
def __rmul__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.multiply, other)
|
||||
|
||||
def __truediv__(self, other: Any) -> Tensor:
|
||||
return self._binary(mx.divide, other, "div")
|
||||
|
||||
def __rtruediv__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.divide, other)
|
||||
|
||||
def __pow__(self, other: Any) -> Tensor:
|
||||
return self.pow(other)
|
||||
|
||||
def __rpow__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.power, other)
|
||||
|
||||
def __matmul__(self, other: Any) -> Tensor:
|
||||
return self._binary(mx.matmul, other, "matmul")
|
||||
|
||||
def __rmatmul__(self, other: Any) -> Tensor:
|
||||
return self._reverse(mx.matmul, other)
|
||||
|
||||
def __neg__(self) -> Tensor:
|
||||
return wrap(mx.negative(self))
|
||||
|
||||
def __eq__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.equal(self, other))
|
||||
|
||||
def __ne__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.not_equal(self, other))
|
||||
|
||||
def __lt__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.less(self, other))
|
||||
|
||||
def __le__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.less_equal(self, other))
|
||||
|
||||
def __gt__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.greater(self, other))
|
||||
|
||||
def __ge__(self, other: Any) -> Tensor:
|
||||
return wrap(mx.greater_equal(self, other))
|
||||
|
||||
|
||||
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
|
||||
mx.array.mean = _torch_mean
|
||||
mx.array.var = _torch_var
|
||||
mx.array.size = _size
|
||||
mx.array.chunk = _chunk
|
||||
mx.array.to = _to
|
||||
mx.array.repeat_interleave = _repeat_interleave
|
||||
mx.array.masked_fill = _masked_fill
|
||||
mx.array.expand = _expand
|
||||
mx.array.any = _torch_any
|
||||
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
|
||||
def _mask_indices(mask: mx.array) -> mx.array:
|
||||
flat = _native_reshape(mask, (-1,))
|
||||
count = int(cast(Any, _native_item(mx.sum(flat))))
|
||||
order = mx.argsort(mx.astype(flat, mx.int32))
|
||||
if count == 0:
|
||||
return mx.astype(_native_getitem(order, slice(0, 0)), mx.int64)
|
||||
return mx.astype(_native_getitem(order, slice(-count, None)), mx.int64)
|
||||
@@ -0,0 +1,57 @@
|
||||
import torch as _native
|
||||
|
||||
|
||||
Tensor = _native.Tensor
|
||||
tensor = _native.tensor
|
||||
from_numpy = _native.from_numpy
|
||||
arange = _native.arange
|
||||
randint = _native.randint
|
||||
chunk = _native.chunk
|
||||
transpose = _native.transpose
|
||||
zeros = _native.zeros
|
||||
ones = _native.ones
|
||||
zeros_like = _native.zeros_like
|
||||
ones_like = _native.ones_like
|
||||
cat = _native.cat
|
||||
softmax = _native.softmax
|
||||
triu = _native.triu
|
||||
tril = _native.tril
|
||||
topk = _native.topk
|
||||
unique = _native.unique
|
||||
exp = _native.exp
|
||||
sin = _native.sin
|
||||
cos = _native.cos
|
||||
tanh = _native.tanh
|
||||
sqrt = _native.sqrt
|
||||
matmul = _native.matmul
|
||||
outer = _native.outer
|
||||
pow = _native.pow
|
||||
polar = _native.polar
|
||||
float16 = _native.float16
|
||||
float32 = _native.float32
|
||||
float64 = _native.float64
|
||||
bfloat16 = _native.bfloat16
|
||||
int8 = _native.int8
|
||||
int16 = _native.int16
|
||||
int32 = _native.int32
|
||||
int64 = _native.int64
|
||||
uint8 = _native.uint8
|
||||
bool = _native.bool
|
||||
half = float16
|
||||
float = float32
|
||||
double = float64
|
||||
int = int32
|
||||
long = int64
|
||||
device = _native.device
|
||||
backends = _native.backends
|
||||
cuda = _native.cuda
|
||||
pi = _native.pi
|
||||
|
||||
|
||||
def categorical(logits, dim=-1, num_samples=1):
|
||||
probabilities = _native.softmax(logits, dim=dim)
|
||||
return _native.multinomial(probabilities, num_samples=num_samples)
|
||||
|
||||
|
||||
def fallback(name):
|
||||
return getattr(_native, name)
|
||||
@@ -1,3 +1,7 @@
|
||||
# pyright: reportAssignmentType=false, reportIncompatibleMethodOverride=false, reportRedeclaration=false
|
||||
|
||||
from typing import cast
|
||||
|
||||
from torchmlx._backend import BACKEND, unsupported
|
||||
|
||||
|
||||
@@ -19,26 +23,62 @@ if BACKEND == "torch":
|
||||
|
||||
else:
|
||||
from collections.abc import Iterable, MutableSequence
|
||||
from contextvars import ContextVar
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as _native
|
||||
from mlx.utils import tree_flatten
|
||||
|
||||
class _ParameterTree(dict):
|
||||
from torchmlx._mlx_tensor import Tensor, wrap
|
||||
|
||||
_module_depth = ContextVar("torchmlx_module_depth", default=0)
|
||||
|
||||
class _ParameterIterator:
|
||||
def __init__(self, values, model):
|
||||
super().__init__(values)
|
||||
self.model = model
|
||||
self._iterator = iter(value for _, value in tree_flatten(values))
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
return wrap(next(self._iterator))
|
||||
|
||||
class Module(_native.Module):
|
||||
def __setattr__(self, name, value):
|
||||
super().__setattr__(name, mx.array(value) if isinstance(value, Tensor) else value)
|
||||
|
||||
def __getattribute__(self, name):
|
||||
value = super().__getattribute__(name)
|
||||
return (
|
||||
wrap(value)
|
||||
if isinstance(value, mx.array) and _module_depth.get() == 0
|
||||
else value
|
||||
)
|
||||
|
||||
def __getattr__(self, name):
|
||||
value = super().__getattr__(name)
|
||||
return (
|
||||
wrap(value)
|
||||
if isinstance(value, mx.array) and _module_depth.get() == 0
|
||||
else value
|
||||
)
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
from torchmlx._autograd import abort_forward, begin_forward, end_forward
|
||||
|
||||
tracking = begin_forward(self)
|
||||
args = wrap(args)
|
||||
kwargs = wrap(kwargs)
|
||||
token = _module_depth.set(_module_depth.get() + 1)
|
||||
try:
|
||||
output = self.forward(*args, **kwargs)
|
||||
output = wrap(self.forward(*args, **kwargs))
|
||||
except Exception:
|
||||
if tracking:
|
||||
abort_forward()
|
||||
raise
|
||||
finally:
|
||||
_module_depth.reset(token)
|
||||
if tracking:
|
||||
end_forward(self, args, kwargs, output)
|
||||
return output
|
||||
@@ -49,7 +89,10 @@ else:
|
||||
)
|
||||
|
||||
def parameters(self):
|
||||
return _ParameterTree(super().parameters(), self)
|
||||
return _ParameterIterator(self._native_parameters(), self)
|
||||
|
||||
def _native_parameters(self):
|
||||
return super().parameters()
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
dtype = kwargs.pop("dtype", None)
|
||||
@@ -80,10 +123,10 @@ else:
|
||||
|
||||
def Parameter(data=None, requires_grad=True):
|
||||
if data is None:
|
||||
return mx.array([])
|
||||
return wrap(mx.array([]))
|
||||
if not requires_grad:
|
||||
unsupported("torchmlx.nn.Parameter with requires_grad=False")
|
||||
return data
|
||||
return wrap(data)
|
||||
|
||||
class Linear(Module, _native.Linear):
|
||||
def __init__(self, in_features, out_features, bias=True, device=None, dtype=None):
|
||||
@@ -156,7 +199,7 @@ else:
|
||||
raise ValueError("the MLX backend only accepts device='mps'")
|
||||
_native.LayerNorm.__init__(
|
||||
self,
|
||||
normalized_shape,
|
||||
cast(int, normalized_shape),
|
||||
eps=eps,
|
||||
affine=elementwise_affine,
|
||||
bias=bias,
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from typing import Any, Iterable, Iterator
|
||||
|
||||
from torchmlx import Tensor
|
||||
from . import functional as functional
|
||||
|
||||
class Module:
|
||||
training: bool
|
||||
def __call__(self, *args: Any, **kwargs: Any) -> Any: ...
|
||||
def forward(self, *args: Any, **kwargs: Any) -> Any: ...
|
||||
def parameters(self) -> Iterator[Tensor]: ...
|
||||
def to(self, *args: Any, **kwargs: Any) -> Module: ...
|
||||
def train(self, mode: bool = True) -> Module: ...
|
||||
def eval(self) -> Module: ...
|
||||
def register_buffer(self, name: str, tensor: Tensor, persistent: bool = True) -> None: ...
|
||||
|
||||
def Parameter(data: Tensor | None = None, requires_grad: bool = True) -> Tensor: ...
|
||||
|
||||
class Linear(Module):
|
||||
weight: Tensor
|
||||
bias: Tensor | None
|
||||
def __init__(self, in_features: int, out_features: int, bias: bool = True, device: Any = None, dtype: Any = None) -> None: ...
|
||||
|
||||
class Embedding(Module):
|
||||
weight: Tensor
|
||||
def __init__(self, num_embeddings: int, embedding_dim: int, padding_idx: int | None = None, max_norm: float | None = None, norm_type: float = 2.0, scale_grad_by_freq: bool = False, sparse: bool = False, device: Any = None, dtype: Any = None) -> None: ...
|
||||
|
||||
class LayerNorm(Module):
|
||||
def __init__(self, normalized_shape: int | Iterable[int], eps: float = 1e-5, elementwise_affine: bool = True, bias: bool = True, device: Any = None, dtype: Any = None) -> None: ...
|
||||
|
||||
class Sequential(Module):
|
||||
def __init__(self, *args: Module) -> None: ...
|
||||
def __len__(self) -> int: ...
|
||||
def __getitem__(self, index: int | str) -> Module: ...
|
||||
|
||||
class GELU(Module):
|
||||
def __init__(self, approximate: str = "none") -> None: ...
|
||||
|
||||
class Dropout(Module):
|
||||
def __init__(self, p: float = 0.5, inplace: bool = False) -> None: ...
|
||||
|
||||
class ModuleList(Module):
|
||||
def __init__(self, modules: Iterable[Module] | None = None) -> None: ...
|
||||
def __getitem__(self, index: int) -> Module: ...
|
||||
def __setitem__(self, index: int, module: Module) -> None: ...
|
||||
def __delitem__(self, index: int) -> None: ...
|
||||
def __len__(self) -> int: ...
|
||||
def __iter__(self) -> Iterator[Module]: ...
|
||||
def insert(self, index: int, module: Module) -> None: ...
|
||||
@@ -1,3 +1,5 @@
|
||||
# pyright: reportAssignmentType=false, reportRedeclaration=false
|
||||
|
||||
import math
|
||||
|
||||
from torchmlx._backend import BACKEND, unsupported
|
||||
@@ -16,6 +18,7 @@ if BACKEND == "torch":
|
||||
|
||||
else:
|
||||
import mlx.core as mx
|
||||
from torchmlx._mlx_tensor import wrap
|
||||
|
||||
def cross_entropy(
|
||||
input,
|
||||
@@ -55,7 +58,7 @@ else:
|
||||
label_smoothing,
|
||||
)
|
||||
return register_loss(
|
||||
loss,
|
||||
wrap(loss),
|
||||
original_input,
|
||||
rebuild,
|
||||
operands,
|
||||
@@ -125,18 +128,20 @@ else:
|
||||
if scale is None:
|
||||
scale = 1 / math.sqrt(query.shape[-1])
|
||||
mask = "causal" if is_causal else attn_mask
|
||||
return mx.fast.scaled_dot_product_attention(
|
||||
query, key, value, scale=scale, mask=mask
|
||||
return wrap(
|
||||
mx.fast.scaled_dot_product_attention(
|
||||
query, key, value, scale=scale, mask=mask
|
||||
)
|
||||
)
|
||||
|
||||
def silu(input, inplace=False):
|
||||
if inplace:
|
||||
unsupported("torchmlx.nn.functional.silu with inplace=True")
|
||||
return input * mx.sigmoid(input)
|
||||
return wrap(input * mx.sigmoid(input))
|
||||
|
||||
def softmax(input, dim=None, dtype=None):
|
||||
value = input.astype(dtype) if dtype is not None else input
|
||||
return mx.softmax(value, axis=dim)
|
||||
return wrap(mx.softmax(value, axis=dim))
|
||||
|
||||
def __getattr__(name):
|
||||
unsupported(f"torchmlx.nn.functional.{name}")
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
from typing import Any
|
||||
from torchmlx import Tensor
|
||||
|
||||
def cross_entropy(input: Tensor, target: Tensor, weight: Tensor | None = None, size_average: bool | None = None, ignore_index: int = -100, reduce: bool | None = None, reduction: str = "mean", label_smoothing: float = 0.0) -> Tensor: ...
|
||||
def scaled_dot_product_attention(query: Tensor, key: Tensor, value: Tensor, attn_mask: Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, scale: float | None = None, enable_gqa: bool = False) -> Tensor: ...
|
||||
def silu(input: Tensor, inplace: bool = False) -> Tensor: ...
|
||||
def softmax(input: Tensor, dim: int | None = None, dtype: Any = None) -> Tensor: ...
|
||||
@@ -1,3 +1,5 @@
|
||||
# pyright: reportAssignmentType=false, reportRedeclaration=false
|
||||
|
||||
from torchmlx._backend import BACKEND, unsupported
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from typing import Any, Iterable
|
||||
from torchmlx import Tensor
|
||||
|
||||
class AdamW:
|
||||
state: Any
|
||||
def __init__(self, params: Iterable[Tensor], lr: float = 1e-3, betas: tuple[float, float] = (0.9, 0.999), eps: float = 1e-8, weight_decay: float = 1e-2, amsgrad: bool = False, maximize: bool = False, foreach: bool | None = None, capturable: bool = False, differentiable: bool = False, fused: bool | None = None) -> None: ...
|
||||
def zero_grad(self, *args: Any, **kwargs: Any) -> None: ...
|
||||
def step(self, *args: Any, **kwargs: Any) -> None: ...
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
from ._backend import BACKEND
|
||||
|
||||
|
||||
if BACKEND == "torch":
|
||||
import torch
|
||||
|
||||
class Trainer:
|
||||
def __init__(self, model, optimizer, loss_fn, compile=True):
|
||||
self.model = model
|
||||
self.optimizer = optimizer
|
||||
self.loss_fn = loss_fn
|
||||
self._step = torch.compile(self._train_step) if compile else self._train_step
|
||||
|
||||
def _train_step(self, x, y):
|
||||
self.optimizer.zero_grad()
|
||||
loss = self.loss_fn(self.model(x), y)
|
||||
loss.backward()
|
||||
self.optimizer.step()
|
||||
return loss
|
||||
|
||||
def step(self, x, y):
|
||||
return self._step(x, y)
|
||||
|
||||
else:
|
||||
class Trainer:
|
||||
def __init__(self, model, optimizer, loss_fn, compile=True):
|
||||
self.model = model
|
||||
self.optimizer = optimizer
|
||||
self.loss_fn = loss_fn
|
||||
self.compile = compile
|
||||
|
||||
def step(self, x, y):
|
||||
from ._autograd import trainer_step
|
||||
|
||||
return trainer_step(self, x, y)
|
||||
@@ -308,6 +308,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/9e/c9/b2622292ea83fbb4ec318f5b9ab867d0a28ab43c5717bb85b0a5f6b3b0a4/networkx-3.6.1-py3-none-any.whl", hash = "sha256:d47fbf302e7d9cbbb9e2555a0d267983d2aa476bac30e90dfbe5669bd57f3762", size = 2068504 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nodeenv"
|
||||
version = "1.10.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/24/bf/d1bda4f6168e0b2e9e5958945e01910052158313224ada5ce1fb2e1113b8/nodeenv-1.10.0.tar.gz", hash = "sha256:996c191ad80897d076bdfba80a41994c2b47c68e224c542b48feba42ba00f8bb", size = 55611 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/88/b2/d0896bdcdc8d28a7fc5717c305f1a861c26e18c05047949fb371034d98bd/nodeenv-1.10.0-py2.py3-none-any.whl", hash = "sha256:5bb13e3eed2923615535339b3c620e76779af4cb4c6a90deccc9e36b274d3827", size = 23438 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "numpy"
|
||||
version = "2.2.6"
|
||||
@@ -686,6 +695,19 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pyright"
|
||||
version = "1.1.414"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "nodeenv" },
|
||||
{ name = "typing-extensions" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/e1/1b/244c7b710031ada80f27e579ec20d28a2285dfc318fed0339866b1047f12/pyright-1.1.414.tar.gz", hash = "sha256:523c0a97c60da6333234955c277730c9cf4f5bd6d5399e7b7d2b0fc5d3599524", size = 4154638 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/d7/ba/18b6e682ead424ad24bcc134339ae5d1b931cd9ae260540592a058a91279/pyright-1.1.414-py3-none-any.whl", hash = "sha256:2a6b4b3298c9eec174c5ed83bd338de6eee82df2992f3e1930e6199d381be36f", size = 6225049 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pytorchmlx"
|
||||
version = "0.0.1"
|
||||
@@ -698,6 +720,11 @@ dependencies = [
|
||||
{ name = "torch" },
|
||||
]
|
||||
|
||||
[package.dev-dependencies]
|
||||
dev = [
|
||||
{ name = "pyright" },
|
||||
]
|
||||
|
||||
[package.metadata]
|
||||
requires-dist = [
|
||||
{ name = "mlx", specifier = ">=0.32.2,<0.33" },
|
||||
@@ -705,6 +732,9 @@ requires-dist = [
|
||||
{ name = "torch", specifier = ">=2.4,<3" },
|
||||
]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
dev = [{ name = "pyright", specifier = ">=1.1.414" }]
|
||||
|
||||
[[package]]
|
||||
name = "setuptools"
|
||||
version = "84.0.0"
|
||||
|
||||
Reference in new issue
Block a user