clean up torchmlx internals

This commit is contained in:
pj committed 2026-09-20 16:49:44 +05:30
1 parent 4eb86111a8
commit cd140fae7a
19 files changed
+1019 -698

No files matched your search

+10
View File
@@ -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.
+5
View File
@@ -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 |
+7 -3
View File
@@ -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.
+11
View File
@@ -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",
]
+7 -242
View File
@@ -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 def __getattr__(name):
return _fallback(name)
class device:
def __init__(self, value):
if isinstance(value, device):
value = value.type
value = str(value)
if value != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
self.type = value
self.index = None
def __str__(self):
return self.type
def __repr__(self):
return f"device(type={self.type!r})"
def __eq__(self, other):
return str(other) == self.type
class _MPS:
@staticmethod
def is_available():
return _platform.system() == "Darwin" and _platform.machine() == "arm64"
class _CUDA:
@staticmethod
def is_available():
return False
backends = _SimpleNamespace(mps=_MPS())
cuda = _CUDA()
from ._mlx_tensor import install as _install_tensor_methods
_install_tensor_methods(device)
def _check_device(device):
if device is not None and str(device) != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
def tensor(data, dtype=None, device=None, requires_grad=False, pin_memory=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
if pin_memory:
raise RuntimeError("pin_memory is not supported by the MLX backend")
if dtype is None and not isinstance(data, (_native.array, _np.ndarray)):
kind = _np.asarray(data).dtype.kind
if kind in {"i", "u"}:
dtype = int64
elif kind == "b":
dtype = bool
return _native.array(data, dtype=dtype)
def from_numpy(array):
if not isinstance(array, _np.ndarray):
raise TypeError("from_numpy expects a numpy.ndarray")
return _native.array(array)
def arange(start, end=None, step=1, *, dtype=None, device=None):
_check_device(device)
if end is None:
start, end = 0, start
if dtype is None and all(
isinstance(value, _builtins.int) for value in (start, end, step)
):
dtype = int64
return _native.arange(start, end, step, dtype=dtype)
def randint(low, high, size, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
result = _native.random.randint(low, high, shape=size)
return result.astype(dtype or int64)
def chunk(input, chunks, dim=0):
return input.chunk(chunks, dim=dim)
def transpose(input, dim0, dim1):
return input.transpose(dim0, dim1)
def zeros(*size, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
return _native.zeros(shape, dtype=dtype or float32)
def ones(*size, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
return _native.ones(shape, dtype=dtype or float32)
def zeros_like(input, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
return _native.zeros_like(input).astype(dtype or input.dtype)
def ones_like(input, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
return _native.ones_like(input).astype(dtype or input.dtype)
def cat(tensors, dim=0):
return _native.concatenate(tensors, axis=dim)
def softmax(input, dim, dtype=None):
value = input.astype(dtype) if dtype is not None else input
return _native.softmax(value, axis=dim)
def triu(input, diagonal=0):
return _native.triu(input, k=diagonal)
def tril(input, diagonal=0):
return _native.tril(input, k=diagonal)
def topk(input, k, dim=None, largest=True, sorted=True):
axis = -1 if dim is None else dim
indices = _native.argsort(input, axis=axis)
indices = _native.flip(indices, axis=axis) if largest else indices
slices = [slice(None)] * input.ndim
slices[axis] = slice(0, k)
indices = indices[tuple(slices)].astype(int64)
values = _native.take_along_axis(input, indices, axis=axis)
return values, indices
def unique(input, sorted=True, return_inverse=False, return_counts=False, dim=None):
if return_inverse or return_counts or dim is not None:
unsupported("torchmlx.unique with non-default options")
values = _native.sort(input.reshape(-1))
if values.shape[0] < 2:
return values
keep = _native.concatenate(
[_native.array([True]), values[1:] != values[:-1]]
)
return values[keep]
exp = _native.exp
sin = _native.sin
cos = _native.cos
tanh = _native.tanh
sqrt = _native.sqrt
matmul = _native.matmul
outer = _native.outer
def pow(input, exponent):
return _native.power(input, exponent)
def polar(abs, angle):
return abs * _native.exp(_native.array(1j) * angle)
def categorical(logits, dim=-1, num_samples=1):
if num_samples == 1:
return _native.random.categorical(logits, axis=dim)[..., None].astype(int64)
return _native.random.categorical(
logits, axis=dim, num_samples=num_samples
).astype(int64)
def __getattr__(name):
unsupported(f"torchmlx.{name}")
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",
+100
View File
@@ -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
View File
@@ -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,
)
+219
View File
@@ -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
View File
@@ -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
+57
View File
@@ -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)
+50 -7
View File
@@ -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,
+48
View File
@@ -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: ...
+9 -4
View File
@@ -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}")
+7
View File
@@ -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: ...
+2
View File
@@ -1,3 +1,5 @@
# pyright: reportAssignmentType=false, reportRedeclaration=false
from torchmlx._backend import BACKEND, unsupported from torchmlx._backend import BACKEND, unsupported
+8
View File
@@ -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: ...
+1
View File
@@ -0,0 +1 @@
-35
View File
@@ -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)
Generated
+30
View File
@@ -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"