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
+1018
-697
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).
|
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.
|
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.
|
Boolean expert routing and `unique` execute eagerly because their output shapes control Python flow.
|
||||||
|
|
||||||
@@ -18,8 +18,12 @@ loss.backward()
|
|||||||
optimizer.step()
|
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.
|
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]
|
[tool.hatch.build.targets.wheel]
|
||||||
packages = ["src/torchmlx"]
|
packages = ["src/torchmlx"]
|
||||||
|
|
||||||
|
[tool.pyright]
|
||||||
|
typeCheckingMode = "standard"
|
||||||
|
include = ["src", "examples/tinystories-llm/train.py"]
|
||||||
|
venvPath = "."
|
||||||
|
venv = ".venv"
|
||||||
|
|
||||||
|
[dependency-groups]
|
||||||
|
dev = [
|
||||||
|
"pyright>=1.1.414",
|
||||||
|
]
|
||||||
+6
-241
@@ -1,259 +1,25 @@
|
|||||||
import builtins as _builtins
|
from ._backend import BACKEND
|
||||||
import math as _math
|
|
||||||
import platform as _platform
|
|
||||||
from types import SimpleNamespace as _SimpleNamespace
|
|
||||||
|
|
||||||
from ._backend import BACKEND, unsupported
|
|
||||||
|
|
||||||
|
|
||||||
def current_backend():
|
|
||||||
return BACKEND
|
|
||||||
|
|
||||||
|
|
||||||
if BACKEND == "torch":
|
if BACKEND == "torch":
|
||||||
import torch as _native
|
from ._torch_backend import *
|
||||||
|
from ._torch_backend import fallback as _fallback
|
||||||
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)
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
import mlx.core as _native
|
from ._mlx_backend import *
|
||||||
import numpy as _np
|
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):
|
def __getattr__(name):
|
||||||
unsupported(f"torchmlx.{name}")
|
return _fallback(name)
|
||||||
|
|
||||||
|
|
||||||
import importlib as _importlib
|
import importlib as _importlib
|
||||||
|
|
||||||
nn = _importlib.import_module("torchmlx.nn")
|
nn = _importlib.import_module("torchmlx.nn")
|
||||||
optim = _importlib.import_module("torchmlx.optim")
|
optim = _importlib.import_module("torchmlx.optim")
|
||||||
from .trainer import Trainer
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Tensor",
|
"Tensor",
|
||||||
"Trainer",
|
|
||||||
"arange",
|
"arange",
|
||||||
"backends",
|
"backends",
|
||||||
"bfloat16",
|
"bfloat16",
|
||||||
@@ -262,7 +28,6 @@ __all__ = [
|
|||||||
"cat",
|
"cat",
|
||||||
"chunk",
|
"chunk",
|
||||||
"cuda",
|
"cuda",
|
||||||
"current_backend",
|
|
||||||
"device",
|
"device",
|
||||||
"double",
|
"double",
|
||||||
"float",
|
"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: ...
|
||||||
+216
-233
@@ -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.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
|
from ._mlx_tensor import Tensor
|
||||||
_active_optimizer = None
|
|
||||||
_forward_depth = 0
|
|
||||||
_loss_plans = {}
|
|
||||||
_lineage = {}
|
|
||||||
_replayed_losses = {}
|
|
||||||
_suspended = False
|
|
||||||
|
|
||||||
|
|
||||||
|
_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:
|
class _ArraySlot:
|
||||||
def __init__(self, index):
|
index: int
|
||||||
self.index = index
|
|
||||||
|
|
||||||
|
@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:
|
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.model = model
|
||||||
self.optimizer = optimizer
|
self.optimizer = optimizer
|
||||||
self._value_and_grad = nn.value_and_grad(model, objective)
|
self._value_and_grad = nn.value_and_grad(model, objective)
|
||||||
self._compile_enabled = compile
|
self._compile_enabled = compile
|
||||||
self._compiled_backward = None
|
self._compiled_backward: Callable[..., Any] | None = None
|
||||||
self._compiled_fused = None
|
self._compiled_fused: Callable[..., Any] | None = None
|
||||||
self._backward_signatures = set()
|
self._backward_signatures: set[Any] = set()
|
||||||
self._fused_signatures = set()
|
self._fused_signatures: set[Any] = set()
|
||||||
if compile:
|
if compile:
|
||||||
backward_state = [model.state, mx.random.state]
|
backward_state = [model.state, mx.random.state]
|
||||||
fused_state = [model.state, optimizer.state, mx.random.state]
|
fused_state = [model.state, optimizer.state, mx.random.state]
|
||||||
self._compiled_backward = mx.compile(
|
self._compiled_backward = mx.compile(
|
||||||
self._value_and_grad,
|
self._value_and_grad, inputs=backward_state, outputs=backward_state
|
||||||
inputs=backward_state,
|
|
||||||
outputs=backward_state,
|
|
||||||
)
|
)
|
||||||
self._compiled_fused = mx.compile(
|
self._compiled_fused = mx.compile(
|
||||||
self._fused,
|
self._fused, inputs=fused_state, outputs=fused_state
|
||||||
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)
|
loss, gradients = self._value_and_grad(*inputs)
|
||||||
self.optimizer.update(self.model, gradients)
|
self.optimizer.update(self.model, gradients)
|
||||||
return loss
|
return loss
|
||||||
|
|
||||||
def backward(self, *inputs):
|
def backward(self, *inputs: mx.array) -> Any:
|
||||||
if self._compile_enabled:
|
if self._compile_enabled:
|
||||||
|
assert self._compiled_backward is not None
|
||||||
return self._compiled_backward(*inputs)
|
return self._compiled_backward(*inputs)
|
||||||
return self._value_and_grad(*inputs)
|
return self._value_and_grad(*inputs)
|
||||||
|
|
||||||
def fused(self, *inputs):
|
def fused(self, *inputs: mx.array) -> mx.array:
|
||||||
if self._compile_enabled:
|
if self._compile_enabled:
|
||||||
|
assert self._compiled_fused is not None
|
||||||
return self._compiled_fused(*inputs)
|
return self._compiled_fused(*inputs)
|
||||||
return self._fused(*inputs)
|
return self._fused(*inputs)
|
||||||
|
|
||||||
def disable_compile(self):
|
def disable_compile(self) -> None:
|
||||||
self._compile_enabled = False
|
self._compile_enabled = False
|
||||||
|
|
||||||
|
|
||||||
def _partition_arrays(value, arrays):
|
def _partition_arrays(value: Any, arrays: list[mx.array]) -> Any:
|
||||||
if isinstance(value, mx.array):
|
if isinstance(value, mx.array):
|
||||||
slot = _ArraySlot(len(arrays))
|
slot = _ArraySlot(len(arrays))
|
||||||
arrays.append(value)
|
arrays.append(value)
|
||||||
@@ -72,7 +105,7 @@ def _partition_arrays(value, arrays):
|
|||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
def _restore_arrays(template, arrays):
|
def _restore_arrays(template: Any, arrays: tuple[mx.array, ...]) -> Any:
|
||||||
if isinstance(template, _ArraySlot):
|
if isinstance(template, _ArraySlot):
|
||||||
return arrays[template.index]
|
return arrays[template.index]
|
||||||
if isinstance(template, tuple):
|
if isinstance(template, tuple):
|
||||||
@@ -84,7 +117,7 @@ def _restore_arrays(template, arrays):
|
|||||||
return template
|
return template
|
||||||
|
|
||||||
|
|
||||||
def _template_key(template):
|
def _template_key(template: Any) -> tuple[Any, ...]:
|
||||||
if isinstance(template, _ArraySlot):
|
if isinstance(template, _ArraySlot):
|
||||||
return ("array",)
|
return ("array",)
|
||||||
if isinstance(template, tuple):
|
if isinstance(template, tuple):
|
||||||
@@ -102,7 +135,7 @@ def _template_key(template):
|
|||||||
return ("static", type(template).__qualname__, repr(template))
|
return ("static", type(template).__qualname__, repr(template))
|
||||||
|
|
||||||
|
|
||||||
def _copy_tree(value):
|
def _copy_tree(value: Any) -> Any:
|
||||||
if isinstance(value, dict):
|
if isinstance(value, dict):
|
||||||
return {key: _copy_tree(item) for key, item in value.items()}
|
return {key: _copy_tree(item) for key, item in value.items()}
|
||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
@@ -112,211 +145,198 @@ def _copy_tree(value):
|
|||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
def _snapshot_random_state():
|
def _snapshot_random_state() -> list[mx.array]:
|
||||||
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
|
random_state = cast(Any, mx.random.state)
|
||||||
|
state = [key + mx.array(0, dtype=key.dtype) for key in random_state]
|
||||||
mx.eval(state)
|
mx.eval(state)
|
||||||
return state
|
return state
|
||||||
|
|
||||||
|
|
||||||
def _restore_random_state(state):
|
def _restore_random_state(state: Any) -> None:
|
||||||
for current_key, saved_key in zip(mx.random.state, state):
|
for current_key, saved_key in zip(cast(Any, mx.random.state), state):
|
||||||
current_key[...] = saved_key
|
current_key[...] = saved_key
|
||||||
mx.eval(mx.random.state)
|
mx.eval(mx.random.state)
|
||||||
|
|
||||||
|
|
||||||
def _compile_failure(error):
|
def _compile_failure(error: Exception) -> bool:
|
||||||
message = str(error)
|
message = str(error)
|
||||||
return "Attempting to eval an array" in message and (
|
return "Attempting to eval an array" in message and (
|
||||||
"function transformations" in message or "without a primitive" in message
|
"function transformations" in message or "without a primitive" in message
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def begin_forward(model):
|
def begin_forward(model: Any) -> bool:
|
||||||
global _forward_depth
|
optimizer = _active_optimizer.get()
|
||||||
if (
|
if _suspended.get() or optimizer is None or model is not optimizer._model:
|
||||||
_suspended
|
|
||||||
or _active_optimizer is None
|
|
||||||
or model is not _active_optimizer._model
|
|
||||||
):
|
|
||||||
return False
|
return False
|
||||||
if _forward_depth:
|
depth = _forward_depth.get()
|
||||||
_active_optimizer._recursive_forward = True
|
if depth:
|
||||||
_forward_depth += 1
|
optimizer._recursive_forward = True
|
||||||
|
_forward_depth.set(depth + 1)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def note_eager_evaluation():
|
def note_eager_evaluation() -> None:
|
||||||
if _forward_depth and _active_optimizer is not None and not _suspended:
|
optimizer = _active_optimizer.get()
|
||||||
_active_optimizer._eager_forward = True
|
if _forward_depth.get() and optimizer is not None and not _suspended.get():
|
||||||
|
optimizer._eager_forward = True
|
||||||
|
|
||||||
|
|
||||||
def end_forward(model, args, kwargs, output):
|
def end_forward(
|
||||||
global _forward_depth
|
model: Any, args: tuple[Any, ...], kwargs: dict[str, Any], output: Any
|
||||||
_forward_depth -= 1
|
) -> None:
|
||||||
if _forward_depth != 0 or not isinstance(output, mx.array):
|
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
|
return
|
||||||
random_before = _active_optimizer._random_before_forward
|
random_before = optimizer._random_before_forward
|
||||||
stochastic = any(
|
stochastic = any(
|
||||||
current is not previous
|
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,
|
model,
|
||||||
args,
|
args,
|
||||||
kwargs,
|
kwargs,
|
||||||
(),
|
(),
|
||||||
random_before if stochastic else None,
|
random_before if stochastic else None,
|
||||||
not (
|
not (optimizer._recursive_forward or optimizer._eager_forward),
|
||||||
_active_optimizer._recursive_forward
|
|
||||||
or _active_optimizer._eager_forward
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
_lineage[id(output)] = (output, descriptor)
|
|
||||||
|
|
||||||
|
|
||||||
def abort_forward():
|
def abort_forward() -> None:
|
||||||
global _forward_depth
|
depth = max(0, _forward_depth.get() - 1)
|
||||||
_forward_depth -= 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):
|
def activate(optimizer: Any) -> None:
|
||||||
global _active_optimizer, _forward_depth
|
previous = _active_optimizer.get()
|
||||||
_active_optimizer = optimizer
|
if previous is not None and previous is not optimizer:
|
||||||
_forward_depth = 0
|
previous._pending_update = None
|
||||||
_loss_plans.clear()
|
_active_optimizer.set(optimizer)
|
||||||
_lineage.clear()
|
_forward_depth.set(0)
|
||||||
_replayed_losses.clear()
|
optimizer._random_before_forward = tuple(cast(Any, mx.random.state))
|
||||||
optimizer._random_before_forward = tuple(mx.random.state)
|
|
||||||
optimizer._recursive_forward = False
|
optimizer._recursive_forward = False
|
||||||
optimizer._eager_forward = False
|
optimizer._eager_forward = False
|
||||||
|
|
||||||
|
|
||||||
def propagate(source, result, operation, signature, operands=()):
|
def propagate(
|
||||||
if _suspended:
|
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
|
return result
|
||||||
entry = _lineage.get(id(source))
|
lineage: _Lineage = source._lineage
|
||||||
if entry is None or entry[0] is not source:
|
result._lineage = _Lineage(
|
||||||
return result
|
lineage.model,
|
||||||
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
|
lineage.args,
|
||||||
_lineage[id(result)] = (
|
lineage.kwargs,
|
||||||
result,
|
lineage.transforms + ((operation, signature, operands),),
|
||||||
(
|
lineage.random_before,
|
||||||
model,
|
lineage.compile_allowed,
|
||||||
args,
|
|
||||||
kwargs,
|
|
||||||
transforms + ((operation, signature, operands),),
|
|
||||||
random_before,
|
|
||||||
compile_allowed,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _plan_for(
|
def _plan_for(optimizer: Any, loss_plan: _LossPlan) -> _TrainingPlan:
|
||||||
optimizer, model, cache_key, template, transforms, rebuild, compile_allowed
|
plan = optimizer._training_plans.get(loss_plan.cache_key)
|
||||||
):
|
|
||||||
plan = optimizer._training_plans.get(cache_key)
|
|
||||||
if plan is not None:
|
if plan is not None:
|
||||||
return plan
|
return plan
|
||||||
|
|
||||||
def objective(*current_inputs):
|
def objective(*current_inputs: mx.array) -> mx.array:
|
||||||
current_args, current_kwargs, current_transform_operands, current_operands = (
|
current_args, current_kwargs, transform_operands, current_operands = (
|
||||||
_restore_arrays(template, current_inputs)
|
_restore_arrays(loss_plan.template, current_inputs)
|
||||||
)
|
)
|
||||||
output = model(*current_args, **current_kwargs)
|
output = loss_plan.model(*current_args, **current_kwargs)
|
||||||
for (operation, _, _), operation_operands in zip(
|
for (operation, _, _), operands in zip(
|
||||||
transforms, current_transform_operands
|
loss_plan.transforms, transform_operands
|
||||||
):
|
):
|
||||||
output = operation(output, *operation_operands)
|
output = operation(output, *operands)
|
||||||
return rebuild(output, *current_operands)
|
return loss_plan.rebuild(output, *current_operands)
|
||||||
|
|
||||||
plan = _TrainingPlan(
|
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
|
return plan
|
||||||
|
|
||||||
|
|
||||||
def register_loss(loss, loss_input, rebuild, operands, signature):
|
def register_loss(
|
||||||
if _suspended or _active_optimizer is None:
|
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
|
return loss
|
||||||
entry = _lineage.get(id(loss_input))
|
lineage: _Lineage = loss_input._lineage
|
||||||
if entry is None or entry[0] is not loss_input:
|
|
||||||
return loss
|
|
||||||
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
|
|
||||||
context = (
|
context = (
|
||||||
args,
|
lineage.args,
|
||||||
kwargs,
|
lineage.kwargs,
|
||||||
tuple(transform_operands for _, _, transform_operands in transforms),
|
tuple(cast(Any, item)[2] for item in lineage.transforms),
|
||||||
operands,
|
operands,
|
||||||
)
|
)
|
||||||
dynamic_inputs = []
|
dynamic_inputs: list[mx.array] = []
|
||||||
template = _partition_arrays(context, dynamic_inputs)
|
template = _partition_arrays(context, dynamic_inputs)
|
||||||
cache_key = (
|
cache_key = (
|
||||||
getattr(model, "training", None),
|
getattr(lineage.model, "training", None),
|
||||||
_template_key(template),
|
_template_key(template),
|
||||||
tuple(transform_signature for _, transform_signature, _ in transforms),
|
tuple(item[1] for item in lineage.transforms),
|
||||||
signature,
|
signature,
|
||||||
compile_allowed,
|
lineage.compile_allowed,
|
||||||
)
|
)
|
||||||
_loss_plans[id(loss)] = (
|
loss._loss_plan = _LossPlan(
|
||||||
loss,
|
lineage.model,
|
||||||
model,
|
|
||||||
cache_key,
|
cache_key,
|
||||||
tuple(dynamic_inputs),
|
tuple(dynamic_inputs),
|
||||||
template,
|
template,
|
||||||
transforms,
|
lineage.transforms,
|
||||||
rebuild,
|
rebuild,
|
||||||
random_before,
|
lineage.random_before,
|
||||||
compile_allowed,
|
lineage.compile_allowed,
|
||||||
)
|
)
|
||||||
return loss
|
return loss
|
||||||
|
|
||||||
|
|
||||||
def backward(loss):
|
def backward(loss: Tensor) -> None:
|
||||||
if _active_optimizer is None:
|
optimizer = _active_optimizer.get()
|
||||||
|
if optimizer is None:
|
||||||
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
|
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
|
||||||
entry = _loss_plans.get(id(loss))
|
loss_plan = loss._loss_plan
|
||||||
if entry is None or entry[0] is not loss:
|
if loss_plan is None:
|
||||||
|
_clear(optimizer)
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
|
"this MLX loss cannot use backward compatibility; compute one supported loss from one model forward"
|
||||||
)
|
)
|
||||||
(
|
try:
|
||||||
_,
|
plan = _plan_for(optimizer, loss_plan)
|
||||||
model,
|
except Exception:
|
||||||
cache_key,
|
_clear(optimizer)
|
||||||
dynamic_inputs,
|
raise
|
||||||
template,
|
optimizer._pending_update = (
|
||||||
transforms,
|
|
||||||
rebuild,
|
|
||||||
random_before,
|
|
||||||
compile_allowed,
|
|
||||||
) = entry
|
|
||||||
plan = _plan_for(
|
|
||||||
_active_optimizer,
|
|
||||||
model,
|
|
||||||
cache_key,
|
|
||||||
template,
|
|
||||||
transforms,
|
|
||||||
rebuild,
|
|
||||||
compile_allowed,
|
|
||||||
)
|
|
||||||
_active_optimizer._pending_update = (
|
|
||||||
"staged",
|
"staged",
|
||||||
loss,
|
loss,
|
||||||
plan,
|
plan,
|
||||||
dynamic_inputs,
|
loss_plan.dynamic_inputs,
|
||||||
random_before,
|
loss_plan.random_before,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _restore_execution(model, optimizer, parameters, optimizer_state):
|
def _restore_tree(current: Any, saved: Any) -> Any:
|
||||||
model.update(parameters)
|
|
||||||
if optimizer_state is not None:
|
|
||||||
_restore_tree(optimizer.state, optimizer_state)
|
|
||||||
|
|
||||||
|
|
||||||
def _restore_tree(current, saved):
|
|
||||||
if isinstance(current, dict) and isinstance(saved, dict):
|
if isinstance(current, dict) and isinstance(saved, dict):
|
||||||
for key in tuple(current):
|
for key in tuple(current):
|
||||||
if key not in saved:
|
if key not in saved:
|
||||||
@@ -330,16 +350,26 @@ def _restore_tree(current, saved):
|
|||||||
current[key] = value
|
current[key] = value
|
||||||
return current
|
return current
|
||||||
if isinstance(current, list) and isinstance(saved, list):
|
if isinstance(current, list) and isinstance(saved, list):
|
||||||
restored = [
|
restored = [_restore_tree(old, value) for old, value in zip(current, saved)]
|
||||||
_restore_tree(old, value) for old, value in zip(current, saved)
|
|
||||||
]
|
|
||||||
current[:] = restored + saved[len(restored) :]
|
current[:] = restored + saved[len(restored) :]
|
||||||
return current
|
return current
|
||||||
return saved
|
return saved
|
||||||
|
|
||||||
|
|
||||||
def _execute(plan, inputs, fused, random_before, advance_random=False):
|
def _restore_execution(
|
||||||
global _suspended
|
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
|
model = plan.model
|
||||||
optimizer = plan.optimizer
|
optimizer = plan.optimizer
|
||||||
signatures = plan._fused_signatures if fused else plan._backward_signatures
|
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)
|
else getattr(optimizer, "_torchmlx_state", None)
|
||||||
)
|
)
|
||||||
random_after = _snapshot_random_state() if random_before is not None else 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:
|
if random_before is not None:
|
||||||
_restore_random_state(random_before)
|
_restore_random_state(random_before)
|
||||||
_suspended = True
|
token = _suspended.set(True)
|
||||||
try:
|
try:
|
||||||
try:
|
try:
|
||||||
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
|
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 []
|
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
|
||||||
mx.eval(
|
mx.eval(
|
||||||
result,
|
result,
|
||||||
current_parameters,
|
current_parameters,
|
||||||
current_optimizer_state,
|
current_optimizer_state,
|
||||||
mx.random.state
|
mx.random.state if random_before is not None else [],
|
||||||
if random_before is not None or advance_random
|
|
||||||
else [],
|
|
||||||
)
|
)
|
||||||
signatures.add(signature)
|
signatures.add(signature)
|
||||||
if fused:
|
if fused:
|
||||||
@@ -385,19 +412,15 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
|
|||||||
_restore_execution(model, optimizer, parameters, optimizer_state)
|
_restore_execution(model, optimizer, parameters, optimizer_state)
|
||||||
if random_before is not None:
|
if random_before is not None:
|
||||||
_restore_random_state(random_before)
|
_restore_random_state(random_before)
|
||||||
elif random_rollback is not None:
|
|
||||||
_restore_random_state(random_rollback)
|
|
||||||
plan.disable_compile()
|
plan.disable_compile()
|
||||||
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
|
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 []
|
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
|
||||||
mx.eval(
|
mx.eval(
|
||||||
result,
|
result,
|
||||||
current_parameters,
|
current_parameters,
|
||||||
current_optimizer_state,
|
current_optimizer_state,
|
||||||
mx.random.state
|
mx.random.state if random_before is not None else [],
|
||||||
if random_before is not None or advance_random
|
|
||||||
else [],
|
|
||||||
)
|
)
|
||||||
if fused:
|
if fused:
|
||||||
optimizer._torchmlx_parameters = current_parameters
|
optimizer._torchmlx_parameters = current_parameters
|
||||||
@@ -405,52 +428,45 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
|
|||||||
except Exception:
|
except Exception:
|
||||||
if parameters is not None:
|
if parameters is not None:
|
||||||
_restore_execution(model, optimizer, parameters, optimizer_state)
|
_restore_execution(model, optimizer, parameters, optimizer_state)
|
||||||
if random_rollback is not None:
|
|
||||||
_restore_random_state(random_rollback)
|
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
_suspended = False
|
_suspended.reset(token)
|
||||||
if random_after is not None:
|
if random_after is not None:
|
||||||
_restore_random_state(random_after)
|
_restore_random_state(random_after)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _materialize(optimizer):
|
def _materialize(optimizer: Any) -> mx.array:
|
||||||
pending = optimizer._pending_update
|
_, loss, plan, dynamic_inputs, random_before = optimizer._pending_update
|
||||||
_, loss, plan, dynamic_inputs, random_before = pending
|
|
||||||
replayed_loss, gradients = _execute(
|
replayed_loss, gradients = _execute(
|
||||||
plan, dynamic_inputs, fused=False, random_before=random_before
|
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)
|
optimizer._pending_update = ("materialized", loss, plan.model, gradients)
|
||||||
return replayed_loss
|
return replayed_loss
|
||||||
|
|
||||||
|
|
||||||
def replayed_value(value):
|
def replayed_value(value: Tensor) -> mx.array:
|
||||||
if _active_optimizer is not None:
|
optimizer = _active_optimizer.get()
|
||||||
pending = _active_optimizer._pending_update
|
if optimizer is not None:
|
||||||
|
pending = optimizer._pending_update
|
||||||
if pending is not None and pending[0] == "staged" and pending[1] is value:
|
if pending is not None and pending[0] == "staged" and pending[1] is value:
|
||||||
return _materialize(_active_optimizer)
|
return _materialize(optimizer)
|
||||||
entry = _replayed_losses.get(id(value))
|
return value._replayed_value if value._replayed_value is not None else value
|
||||||
if entry is None or entry[0] is not value:
|
|
||||||
return value
|
|
||||||
return entry[1]
|
|
||||||
|
|
||||||
|
|
||||||
def _finish_step(loss, replayed_loss):
|
def _clear(optimizer: Any) -> None:
|
||||||
global _active_optimizer, _forward_depth
|
optimizer._pending_update = None
|
||||||
_replayed_losses[id(loss)] = (loss, replayed_loss)
|
_active_optimizer.set(None)
|
||||||
_active_optimizer._pending_update = None
|
_forward_depth.set(0)
|
||||||
_active_optimizer = None
|
|
||||||
_forward_depth = 0
|
|
||||||
_loss_plans.clear()
|
|
||||||
_lineage.clear()
|
|
||||||
|
|
||||||
|
|
||||||
def step(optimizer):
|
def step(optimizer: Any) -> None:
|
||||||
pending = optimizer._pending_update
|
pending = optimizer._pending_update
|
||||||
if pending is None:
|
if pending is None:
|
||||||
|
_clear(optimizer)
|
||||||
raise RuntimeError("loss.backward() must be called before optimizer.step()")
|
raise RuntimeError("loss.backward() must be called before optimizer.step()")
|
||||||
|
try:
|
||||||
if pending[0] == "staged":
|
if pending[0] == "staged":
|
||||||
_, loss, plan, dynamic_inputs, random_before = pending
|
_, loss, plan, dynamic_inputs, random_before = pending
|
||||||
replayed_loss = _execute(
|
replayed_loss = _execute(
|
||||||
@@ -462,48 +478,15 @@ def step(optimizer):
|
|||||||
optimizer_state = _copy_tree(optimizer._optimizer.state)
|
optimizer_state = _copy_tree(optimizer._optimizer.state)
|
||||||
try:
|
try:
|
||||||
optimizer._compiled_step(gradients)
|
optimizer._compiled_step(gradients)
|
||||||
current_parameters = model.parameters()
|
current_parameters = model._native_parameters()
|
||||||
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
|
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
|
||||||
mx.eval(current_parameters, current_optimizer_state)
|
mx.eval(current_parameters, current_optimizer_state)
|
||||||
except Exception:
|
except Exception:
|
||||||
_restore_execution(
|
_restore_execution(model, optimizer._optimizer, parameters, optimizer_state)
|
||||||
model, optimizer._optimizer, parameters, optimizer_state
|
|
||||||
)
|
|
||||||
raise
|
raise
|
||||||
optimizer._optimizer._torchmlx_parameters = current_parameters
|
optimizer._optimizer._torchmlx_parameters = current_parameters
|
||||||
optimizer._optimizer._torchmlx_state = current_optimizer_state
|
optimizer._optimizer._torchmlx_state = current_optimizer_state
|
||||||
replayed_loss = _replayed_losses[id(loss)][1]
|
replayed_loss = loss._replayed_value
|
||||||
_finish_step(loss, replayed_loss)
|
loss._replayed_value = replayed_loss
|
||||||
|
finally:
|
||||||
|
_clear(optimizer)
|
||||||
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,
|
|
||||||
)
|
|
||||||
@@ -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}")
|
||||||
+232
-174
@@ -1,136 +1,157 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any, Callable, cast
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
_transpose = mx.array.transpose
|
_native_getitem = mx.array.__getitem__
|
||||||
_reshape = mx.array.reshape
|
_native_setitem = mx.array.__setitem__
|
||||||
_squeeze = mx.array.squeeze
|
_native_item = mx.array.item
|
||||||
_mean = mx.array.mean
|
_native_reshape = mx.array.reshape
|
||||||
_var = mx.array.var
|
_native_squeeze = mx.array.squeeze
|
||||||
_any = mx.array.any
|
_native_transpose = mx.array.transpose
|
||||||
_getitem = mx.array.__getitem__
|
|
||||||
_setitem = mx.array.__setitem__
|
|
||||||
_item = mx.array.item
|
|
||||||
|
|
||||||
|
|
||||||
def _torch_transpose(self, dim0=None, dim1=None):
|
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:
|
if dim1 is None:
|
||||||
if dim0 is None:
|
axes = list(reversed(range(self.ndim))) if dim0 is None else dim0
|
||||||
axes = list(reversed(range(self.ndim)))
|
|
||||||
else:
|
|
||||||
axes = dim0
|
|
||||||
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]
|
||||||
result = _transpose(self, axes)
|
return self._result(
|
||||||
from ._autograd import propagate
|
_native_transpose(self, axes),
|
||||||
|
lambda value: _native_transpose(value, axes),
|
||||||
return propagate(
|
|
||||||
self,
|
|
||||||
result,
|
|
||||||
lambda value: _transpose(value, axes),
|
|
||||||
("transpose", tuple(axes)),
|
("transpose", tuple(axes)),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def view(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
|
return self.reshape(*shape)
|
||||||
|
|
||||||
def _view(self, *shape):
|
def reshape(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
|
||||||
shape = shape[0]
|
shape = tuple(shape[0])
|
||||||
return _torch_reshape(self, shape)
|
return self._result(
|
||||||
|
_native_reshape(self, shape),
|
||||||
|
lambda value: _native_reshape(value, 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)),
|
("reshape", tuple(shape)),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def unsqueeze(self, dim: int) -> Tensor:
|
||||||
def _unsqueeze(self, dim):
|
return self._result(
|
||||||
result = mx.expand_dims(self, axis=dim)
|
mx.expand_dims(self, axis=dim),
|
||||||
from ._autograd import propagate
|
|
||||||
|
|
||||||
return propagate(
|
|
||||||
self,
|
|
||||||
result,
|
|
||||||
lambda value: mx.expand_dims(value, axis=dim),
|
lambda value: mx.expand_dims(value, axis=dim),
|
||||||
("unsqueeze", dim),
|
("unsqueeze", dim),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def squeeze(self, dim: int | None = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
def _torch_squeeze(self, dim=None):
|
return self._result(
|
||||||
result = _squeeze(self, axis=dim)
|
_native_squeeze(self, axis=dim),
|
||||||
from ._autograd import propagate
|
lambda value: _native_squeeze(value, axis=dim),
|
||||||
|
|
||||||
return propagate(
|
|
||||||
self,
|
|
||||||
result,
|
|
||||||
lambda value: _squeeze(value, axis=dim),
|
|
||||||
("squeeze", dim),
|
("squeeze", dim),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def flatten(self, start_dim: int = 0, end_dim: int = -1) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
def _flatten(self, start_dim=0, end_dim=-1):
|
|
||||||
if end_dim < 0:
|
if end_dim < 0:
|
||||||
end_dim += self.ndim
|
end_dim += self.ndim
|
||||||
flattened = 1
|
flattened = 1
|
||||||
for dimension in self.shape[start_dim : end_dim + 1]:
|
for dimension in self.shape[start_dim : end_dim + 1]:
|
||||||
flattened *= dimension
|
flattened *= dimension
|
||||||
shape = self.shape[:start_dim] + (flattened,) + self.shape[end_dim + 1 :]
|
shape = self.shape[:start_dim] + (flattened,) + self.shape[end_dim + 1 :]
|
||||||
return _torch_reshape(self, shape)
|
return self.reshape(shape)
|
||||||
|
|
||||||
|
def float(self) -> Tensor:
|
||||||
def _float(self):
|
|
||||||
return self.astype(mx.float32)
|
return self.astype(mx.float32)
|
||||||
|
|
||||||
|
def bool(self) -> Tensor:
|
||||||
def _bool(self):
|
|
||||||
return self.astype(mx.bool_)
|
return self.astype(mx.bool_)
|
||||||
|
|
||||||
|
def pow(self, exponent: Any) -> Tensor:
|
||||||
|
return self._binary(mx.power, exponent, "pow")
|
||||||
|
|
||||||
def _pow(self, exponent):
|
def mean( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
return mx.power(self, exponent)
|
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 var( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
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,
|
self,
|
||||||
dim=None,
|
dim: Any = None,
|
||||||
unbiased=True,
|
unbiased: bool = True,
|
||||||
keepdim=False,
|
keepdim: bool = False,
|
||||||
*,
|
*,
|
||||||
correction=None,
|
correction: int | None = None,
|
||||||
):
|
) -> Tensor:
|
||||||
ddof = int(unbiased) if correction is None else correction
|
ddof = int(unbiased) if correction is None else correction
|
||||||
return _var(self, axis=dim, keepdims=keepdim, ddof=ddof)
|
return wrap(mx.var(self, axis=dim, keepdims=keepdim, ddof=ddof))
|
||||||
|
|
||||||
|
def size(self, dim: int | None = None) -> Any: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
def _size(self, dim=None):
|
|
||||||
return self.shape if dim is None else self.shape[dim]
|
return self.shape if dim is None else self.shape[dim]
|
||||||
|
|
||||||
|
def chunk(self, chunks: int, dim: int = 0) -> tuple[Tensor, ...]:
|
||||||
def _chunk(self, chunks, dim=0):
|
|
||||||
if chunks <= 0:
|
if chunks <= 0:
|
||||||
raise ValueError("chunks must be greater than 0")
|
raise ValueError("chunks must be greater than 0")
|
||||||
length = self.shape[dim]
|
length = self.shape[dim]
|
||||||
if length == 0:
|
if length == 0:
|
||||||
return tuple(mx.split(self, chunks, axis=dim))
|
return wrap(tuple(mx.split(self, chunks, axis=dim)))
|
||||||
chunk_size = (length + chunks - 1) // chunks
|
chunk_size = (length + chunks - 1) // chunks
|
||||||
indices = list(range(chunk_size, length, chunk_size))
|
indices = list(range(chunk_size, length, chunk_size))
|
||||||
return tuple(mx.split(self, indices, axis=dim))
|
return wrap(tuple(mx.split(self, indices, axis=dim)))
|
||||||
|
|
||||||
|
def to(
|
||||||
def _to(self, *args, dtype=None, device=None, **kwargs):
|
self,
|
||||||
|
*args: Any,
|
||||||
|
dtype: Any = None,
|
||||||
|
device: Any = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Tensor:
|
||||||
if kwargs:
|
if kwargs:
|
||||||
name = next(iter(kwargs))
|
name = next(iter(kwargs))
|
||||||
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
|
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
|
||||||
@@ -145,131 +166,168 @@ def _to(self, *args, dtype=None, device=None, **kwargs):
|
|||||||
raise ValueError("the MLX backend only accepts 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
|
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 _repeat_interleave(self, repeats, dim=None):
|
def masked_fill(self, mask: mx.array, value: Any) -> Tensor:
|
||||||
return mx.repeat(self, repeats, axis=dim)
|
return wrap(mx.where(mask, mx.array(value, dtype=self.dtype), self))
|
||||||
|
|
||||||
|
def expand(self, *sizes: Any) -> Tensor:
|
||||||
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)):
|
if len(sizes) == 1 and isinstance(sizes[0], (tuple, list)):
|
||||||
sizes = tuple(sizes[0])
|
sizes = tuple(sizes[0])
|
||||||
if len(sizes) < self.ndim:
|
if len(sizes) < self.ndim:
|
||||||
raise ValueError("expanded size must have at least as many dimensions as the tensor")
|
raise ValueError("expanded size must have at least as many dimensions as the tensor")
|
||||||
source = (1,) * (len(sizes) - self.ndim) + self.shape
|
source = (1,) * (len(sizes) - self.ndim) + self.shape
|
||||||
target = tuple(current if requested == -1 else requested for requested, current in zip(sizes, source))
|
target = tuple(
|
||||||
return mx.broadcast_to(self.reshape(source), target)
|
current if requested == -1 else requested
|
||||||
|
for requested, current in zip(sizes, source)
|
||||||
|
|
||||||
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"
|
|
||||||
)
|
)
|
||||||
|
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
|
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, **kwargs):
|
def backward(self, *args: Any, **kwargs: Any) -> None:
|
||||||
if args or kwargs:
|
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")
|
raise TypeError("MLX backward compatibility does not accept arguments")
|
||||||
from ._autograd import backward
|
from ._autograd import backward
|
||||||
|
|
||||||
backward(self)
|
backward(self)
|
||||||
|
|
||||||
|
def item(self) -> Any:
|
||||||
def _torch_item(self):
|
|
||||||
from ._autograd import note_eager_evaluation, replayed_value
|
from ._autograd import note_eager_evaluation, replayed_value
|
||||||
|
|
||||||
note_eager_evaluation()
|
note_eager_evaluation()
|
||||||
return _item(replayed_value(self))
|
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 _mask_indices(mask):
|
def __getitem__(self, key: Any) -> Tensor:
|
||||||
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 _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:
|
||||||
result = _getitem(self.reshape(-1), indices)
|
result = _native_getitem(_native_reshape(self, (-1,)), indices)
|
||||||
operation = lambda value, current_indices: _getitem(
|
operation = lambda value, current: _native_getitem(
|
||||||
_reshape(value, (-1,)), current_indices
|
_native_reshape(value, (-1,)), current
|
||||||
)
|
)
|
||||||
signature = ("getitem_bool", "flat")
|
signature = ("getitem_bool", "flat")
|
||||||
else:
|
else:
|
||||||
result = _getitem(self, indices)
|
result = _native_getitem(self, indices)
|
||||||
operation = lambda value, current_indices: _getitem(
|
operation = lambda value, current: _native_getitem(value, current)
|
||||||
value, current_indices
|
|
||||||
)
|
|
||||||
signature = ("getitem_bool", "first_axis")
|
signature = ("getitem_bool", "first_axis")
|
||||||
operands = (indices,)
|
return self._result(result, operation, signature, (indices,))
|
||||||
elif isinstance(key, mx.array):
|
if isinstance(key, mx.array):
|
||||||
result = _getitem(self, key)
|
return self._result(
|
||||||
operation = lambda value, current_key: _getitem(value, current_key)
|
_native_getitem(self, key),
|
||||||
signature = ("getitem_array",)
|
lambda value, current: _native_getitem(value, current),
|
||||||
operands = (key,)
|
("getitem_array",),
|
||||||
else:
|
(key,),
|
||||||
result = _getitem(self, key)
|
)
|
||||||
operation = lambda value: _getitem(value, key)
|
return self._result(
|
||||||
signature = ("getitem", type(key).__qualname__, repr(key))
|
_native_getitem(self, key),
|
||||||
operands = ()
|
lambda value: _native_getitem(value, key),
|
||||||
from ._autograd import propagate
|
("getitem", type(key).__qualname__, repr(key)),
|
||||||
|
)
|
||||||
|
|
||||||
return propagate(self, result, operation, signature, operands)
|
def __setitem__(self, key: Any, value: Any) -> None:
|
||||||
|
|
||||||
|
|
||||||
def _torch_setitem(self, key, value):
|
|
||||||
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:
|
||||||
flat = self.reshape(-1)
|
_native_setitem(_native_reshape(self, (-1,)), indices, value)
|
||||||
_setitem(flat, indices, value)
|
|
||||||
return
|
return
|
||||||
_setitem(self, indices, value)
|
_native_setitem(self, indices, value)
|
||||||
return
|
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):
|
def _mask_indices(mask: mx.array) -> mx.array:
|
||||||
mx.array.transpose = _torch_transpose
|
flat = _native_reshape(mask, (-1,))
|
||||||
mx.array.reshape = _torch_reshape
|
count = int(cast(Any, _native_item(mx.sum(flat))))
|
||||||
mx.array.view = _view
|
order = mx.argsort(mx.astype(flat, mx.int32))
|
||||||
mx.array.unsqueeze = _unsqueeze
|
if count == 0:
|
||||||
mx.array.squeeze = _torch_squeeze
|
return mx.astype(_native_getitem(order, slice(0, 0)), mx.int64)
|
||||||
mx.array.flatten = _flatten
|
return mx.astype(_native_getitem(order, slice(-count, None)), mx.int64)
|
||||||
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
|
|
||||||
@@ -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
|
from torchmlx._backend import BACKEND, unsupported
|
||||||
|
|
||||||
|
|
||||||
@@ -19,26 +23,62 @@ if BACKEND == "torch":
|
|||||||
|
|
||||||
else:
|
else:
|
||||||
from collections.abc import Iterable, MutableSequence
|
from collections.abc import Iterable, MutableSequence
|
||||||
|
from contextvars import ContextVar
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as _native
|
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):
|
def __init__(self, values, model):
|
||||||
super().__init__(values)
|
|
||||||
self.model = model
|
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):
|
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):
|
def __call__(self, *args, **kwargs):
|
||||||
from torchmlx._autograd import abort_forward, begin_forward, end_forward
|
from torchmlx._autograd import abort_forward, begin_forward, end_forward
|
||||||
|
|
||||||
tracking = begin_forward(self)
|
tracking = begin_forward(self)
|
||||||
|
args = wrap(args)
|
||||||
|
kwargs = wrap(kwargs)
|
||||||
|
token = _module_depth.set(_module_depth.get() + 1)
|
||||||
try:
|
try:
|
||||||
output = self.forward(*args, **kwargs)
|
output = wrap(self.forward(*args, **kwargs))
|
||||||
except Exception:
|
except Exception:
|
||||||
if tracking:
|
if tracking:
|
||||||
abort_forward()
|
abort_forward()
|
||||||
raise
|
raise
|
||||||
|
finally:
|
||||||
|
_module_depth.reset(token)
|
||||||
if tracking:
|
if tracking:
|
||||||
end_forward(self, args, kwargs, output)
|
end_forward(self, args, kwargs, output)
|
||||||
return output
|
return output
|
||||||
@@ -49,7 +89,10 @@ else:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def parameters(self):
|
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):
|
def to(self, *args, **kwargs):
|
||||||
dtype = kwargs.pop("dtype", None)
|
dtype = kwargs.pop("dtype", None)
|
||||||
@@ -80,10 +123,10 @@ else:
|
|||||||
|
|
||||||
def Parameter(data=None, requires_grad=True):
|
def Parameter(data=None, requires_grad=True):
|
||||||
if data is None:
|
if data is None:
|
||||||
return mx.array([])
|
return wrap(mx.array([]))
|
||||||
if not requires_grad:
|
if not requires_grad:
|
||||||
unsupported("torchmlx.nn.Parameter with requires_grad=False")
|
unsupported("torchmlx.nn.Parameter with requires_grad=False")
|
||||||
return data
|
return wrap(data)
|
||||||
|
|
||||||
class Linear(Module, _native.Linear):
|
class Linear(Module, _native.Linear):
|
||||||
def __init__(self, in_features, out_features, bias=True, device=None, dtype=None):
|
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'")
|
raise ValueError("the MLX backend only accepts device='mps'")
|
||||||
_native.LayerNorm.__init__(
|
_native.LayerNorm.__init__(
|
||||||
self,
|
self,
|
||||||
normalized_shape,
|
cast(int, normalized_shape),
|
||||||
eps=eps,
|
eps=eps,
|
||||||
affine=elementwise_affine,
|
affine=elementwise_affine,
|
||||||
bias=bias,
|
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
|
import math
|
||||||
|
|
||||||
from torchmlx._backend import BACKEND, unsupported
|
from torchmlx._backend import BACKEND, unsupported
|
||||||
@@ -16,6 +18,7 @@ if BACKEND == "torch":
|
|||||||
|
|
||||||
else:
|
else:
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
from torchmlx._mlx_tensor import wrap
|
||||||
|
|
||||||
def cross_entropy(
|
def cross_entropy(
|
||||||
input,
|
input,
|
||||||
@@ -55,7 +58,7 @@ else:
|
|||||||
label_smoothing,
|
label_smoothing,
|
||||||
)
|
)
|
||||||
return register_loss(
|
return register_loss(
|
||||||
loss,
|
wrap(loss),
|
||||||
original_input,
|
original_input,
|
||||||
rebuild,
|
rebuild,
|
||||||
operands,
|
operands,
|
||||||
@@ -125,18 +128,20 @@ else:
|
|||||||
if scale is None:
|
if scale is None:
|
||||||
scale = 1 / math.sqrt(query.shape[-1])
|
scale = 1 / math.sqrt(query.shape[-1])
|
||||||
mask = "causal" if is_causal else attn_mask
|
mask = "causal" if is_causal else attn_mask
|
||||||
return mx.fast.scaled_dot_product_attention(
|
return wrap(
|
||||||
|
mx.fast.scaled_dot_product_attention(
|
||||||
query, key, value, scale=scale, mask=mask
|
query, key, value, scale=scale, mask=mask
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def silu(input, inplace=False):
|
def silu(input, inplace=False):
|
||||||
if inplace:
|
if inplace:
|
||||||
unsupported("torchmlx.nn.functional.silu with inplace=True")
|
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):
|
def softmax(input, dim=None, dtype=None):
|
||||||
value = input.astype(dtype) if dtype is not None else input
|
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):
|
def __getattr__(name):
|
||||||
unsupported(f"torchmlx.nn.functional.{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
|
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 },
|
{ 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]]
|
[[package]]
|
||||||
name = "numpy"
|
name = "numpy"
|
||||||
version = "2.2.6"
|
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 },
|
{ 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]]
|
[[package]]
|
||||||
name = "pytorchmlx"
|
name = "pytorchmlx"
|
||||||
version = "0.0.1"
|
version = "0.0.1"
|
||||||
@@ -698,6 +720,11 @@ dependencies = [
|
|||||||
{ name = "torch" },
|
{ name = "torch" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[package.dev-dependencies]
|
||||||
|
dev = [
|
||||||
|
{ name = "pyright" },
|
||||||
|
]
|
||||||
|
|
||||||
[package.metadata]
|
[package.metadata]
|
||||||
requires-dist = [
|
requires-dist = [
|
||||||
{ name = "mlx", specifier = ">=0.32.2,<0.33" },
|
{ name = "mlx", specifier = ">=0.32.2,<0.33" },
|
||||||
@@ -705,6 +732,9 @@ requires-dist = [
|
|||||||
{ name = "torch", specifier = ">=2.4,<3" },
|
{ name = "torch", specifier = ">=2.4,<3" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[package.metadata.requires-dev]
|
||||||
|
dev = [{ name = "pyright", specifier = ">=1.1.414" }]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "setuptools"
|
name = "setuptools"
|
||||||
version = "84.0.0"
|
version = "84.0.0"
|
||||||
|
|||||||
Reference in new issue
Block a user