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
+1115 -794

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).
| 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.
MLX arrays remain native arrays. Torch-style tensor methods are installed on the native array type for the supported subset.
On MLX, TorchMLX tensors are an internal MLX array subclass. Native MLX arrays and their methods are not modified. Model parameters remain native MLX arrays internally and are presented as TorchMLX tensors through the PyTorch-shaped interface.
Boolean expert routing and `unique` execute eagerly because their output shapes control Python flow.
@@ -18,8 +18,12 @@ loss.backward()
optimizer.step()
```
MLX implements this sequence by recording the outer model call and replaying the forward, backward, and AdamW update inside one cached compiled graph during `step`. Inputs and loss operands remain dynamic, so batches are not captured as constants. Random state is restored for the replay so dropout uses the same mask. Reading `loss.item()` before `step` materializes gradients through a separate compiled path. Models with eager data-dependent operations fall back to an uncompiled replay. Unrecorded loss expressions, gradient hooks, parameter `.grad`, higher-order gradients, and multiple-forward losses remain unsupported.
MLX implements this sequence by recording the outer model call and replaying the forward, backward, and AdamW update inside one cached compiled graph during `step`. Inputs and loss operands remain dynamic, so batches are not captured as constants. Random state is restored for the replay so dropout uses the same mask. Reading `loss.item()` after `backward` and before `step` materializes gradients through a separate compiled path. Models with eager data-dependent operations fall back to an uncompiled replay.
Set `TORCHMLX_BACKEND=torch` before import to use native PyTorch for unsupported programs. TorchMLX never changes backend during an operation.
The supported subset includes tensor construction, dtype conversion, indexing, arithmetic, comparisons, reshape and view operations, transpose, squeeze and unsqueeze, flatten, expand, chunk, masking, common transformer math, embeddings, linear layers, layer normalization, dropout, GELU, causal attention, cross entropy, and AdamW with its default feature set.
Unrecorded loss expressions, losses combining multiple model forwards, gradient hooks, parameter `.grad`, higher-order gradients, serialization, additional optimizers, and additional losses remain unsupported on MLX. Use native PyTorch when a program needs behavior outside this subset.
TorchMLX selects MLX on Apple silicon and PyTorch elsewhere. `TORCHMLX_BACKEND=torch` and `TORCHMLX_BACKEND=mlx` are internal validation overrides. TorchMLX never changes backend during an operation.
The referenced OpenArch model files contain source errors independent of TorchMLX, including invalid constructor calls and undefined attributes. Correct those errors before using either backend.
+11
View File
@@ -20,3 +20,14 @@ build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["src/torchmlx"]
[tool.pyright]
typeCheckingMode = "standard"
include = ["src", "examples/tinystories-llm/train.py"]
venvPath = "."
venv = ".venv"
[dependency-groups]
dev = [
"pyright>=1.1.414",
]
+7 -242
View File
@@ -1,259 +1,25 @@
import builtins as _builtins
import math as _math
import platform as _platform
from types import SimpleNamespace as _SimpleNamespace
from ._backend import BACKEND, unsupported
def current_backend():
return BACKEND
from ._backend import BACKEND
if BACKEND == "torch":
import torch as _native
Tensor = _native.Tensor
tensor = _native.tensor
from_numpy = _native.from_numpy
arange = _native.arange
randint = _native.randint
chunk = _native.chunk
transpose = _native.transpose
float16 = _native.float16
float32 = _native.float32
float64 = _native.float64
bfloat16 = _native.bfloat16
int8 = _native.int8
int16 = _native.int16
int32 = _native.int32
int64 = _native.int64
uint8 = _native.uint8
bool = _native.bool
half = float16
float = float32
double = float64
int = int32
long = int64
def categorical(logits, dim=-1, num_samples=1):
probabilities = _native.softmax(logits, dim=dim)
return _native.multinomial(probabilities, num_samples=num_samples)
def __getattr__(name):
return getattr(_native, name)
from ._torch_backend import *
from ._torch_backend import fallback as _fallback
else:
import mlx.core as _native
import numpy as _np
from ._mlx_backend import *
from ._mlx_backend import fallback as _fallback
Tensor = _native.array
float16 = _native.float16
float32 = _native.float32
float64 = _native.float64
bfloat16 = _native.bfloat16
int8 = _native.int8
int16 = _native.int16
int32 = _native.int32
int64 = _native.int64
uint8 = _native.uint8
bool = _native.bool_
half = float16
float = float32
double = float64
int = int32
long = int64
pi = _math.pi
class device:
def __init__(self, value):
if isinstance(value, device):
value = value.type
value = str(value)
if value != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
self.type = value
self.index = None
def __str__(self):
return self.type
def __repr__(self):
return f"device(type={self.type!r})"
def __eq__(self, other):
return str(other) == self.type
class _MPS:
@staticmethod
def is_available():
return _platform.system() == "Darwin" and _platform.machine() == "arm64"
class _CUDA:
@staticmethod
def is_available():
return False
backends = _SimpleNamespace(mps=_MPS())
cuda = _CUDA()
from ._mlx_tensor import install as _install_tensor_methods
_install_tensor_methods(device)
def _check_device(device):
if device is not None and str(device) != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
def tensor(data, dtype=None, device=None, requires_grad=False, pin_memory=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
if pin_memory:
raise RuntimeError("pin_memory is not supported by the MLX backend")
if dtype is None and not isinstance(data, (_native.array, _np.ndarray)):
kind = _np.asarray(data).dtype.kind
if kind in {"i", "u"}:
dtype = int64
elif kind == "b":
dtype = bool
return _native.array(data, dtype=dtype)
def from_numpy(array):
if not isinstance(array, _np.ndarray):
raise TypeError("from_numpy expects a numpy.ndarray")
return _native.array(array)
def arange(start, end=None, step=1, *, dtype=None, device=None):
_check_device(device)
if end is None:
start, end = 0, start
if dtype is None and all(
isinstance(value, _builtins.int) for value in (start, end, step)
):
dtype = int64
return _native.arange(start, end, step, dtype=dtype)
def randint(low, high, size, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
result = _native.random.randint(low, high, shape=size)
return result.astype(dtype or int64)
def chunk(input, chunks, dim=0):
return input.chunk(chunks, dim=dim)
def transpose(input, dim0, dim1):
return input.transpose(dim0, dim1)
def zeros(*size, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
return _native.zeros(shape, dtype=dtype or float32)
def ones(*size, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
shape = size[0] if len(size) == 1 and isinstance(size[0], (tuple, list)) else size
return _native.ones(shape, dtype=dtype or float32)
def zeros_like(input, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
return _native.zeros_like(input).astype(dtype or input.dtype)
def ones_like(input, *, dtype=None, device=None, requires_grad=False):
_check_device(device)
if requires_grad:
raise RuntimeError(
"requires_grad is not supported by the MLX backend; use torchmlx.Trainer"
)
return _native.ones_like(input).astype(dtype or input.dtype)
def cat(tensors, dim=0):
return _native.concatenate(tensors, axis=dim)
def softmax(input, dim, dtype=None):
value = input.astype(dtype) if dtype is not None else input
return _native.softmax(value, axis=dim)
def triu(input, diagonal=0):
return _native.triu(input, k=diagonal)
def tril(input, diagonal=0):
return _native.tril(input, k=diagonal)
def topk(input, k, dim=None, largest=True, sorted=True):
axis = -1 if dim is None else dim
indices = _native.argsort(input, axis=axis)
indices = _native.flip(indices, axis=axis) if largest else indices
slices = [slice(None)] * input.ndim
slices[axis] = slice(0, k)
indices = indices[tuple(slices)].astype(int64)
values = _native.take_along_axis(input, indices, axis=axis)
return values, indices
def unique(input, sorted=True, return_inverse=False, return_counts=False, dim=None):
if return_inverse or return_counts or dim is not None:
unsupported("torchmlx.unique with non-default options")
values = _native.sort(input.reshape(-1))
if values.shape[0] < 2:
return values
keep = _native.concatenate(
[_native.array([True]), values[1:] != values[:-1]]
)
return values[keep]
exp = _native.exp
sin = _native.sin
cos = _native.cos
tanh = _native.tanh
sqrt = _native.sqrt
matmul = _native.matmul
outer = _native.outer
def pow(input, exponent):
return _native.power(input, exponent)
def polar(abs, angle):
return abs * _native.exp(_native.array(1j) * angle)
def categorical(logits, dim=-1, num_samples=1):
if num_samples == 1:
return _native.random.categorical(logits, axis=dim)[..., None].astype(int64)
return _native.random.categorical(
logits, axis=dim, num_samples=num_samples
).astype(int64)
def __getattr__(name):
unsupported(f"torchmlx.{name}")
def __getattr__(name):
return _fallback(name)
import importlib as _importlib
nn = _importlib.import_module("torchmlx.nn")
optim = _importlib.import_module("torchmlx.optim")
from .trainer import Trainer
__all__ = [
"Tensor",
"Trainer",
"arange",
"backends",
"bfloat16",
@@ -262,7 +28,6 @@ __all__ = [
"cat",
"chunk",
"cuda",
"current_backend",
"device",
"double",
"float",
+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: ...
+232 -249
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.nn as nn
_active_optimizer = None
_forward_depth = 0
_loss_plans = {}
_lineage = {}
_replayed_losses = {}
_suspended = False
from ._mlx_tensor import Tensor
_active_optimizer: ContextVar[Any | None] = ContextVar(
"torchmlx_active_optimizer", default=None
)
_forward_depth: ContextVar[int] = ContextVar("torchmlx_forward_depth", default=0)
_suspended: ContextVar[bool] = ContextVar("torchmlx_suspended", default=False)
@dataclass(frozen=True)
class _ArraySlot:
def __init__(self, index):
self.index = index
index: int
@dataclass(frozen=True)
class _Lineage:
model: Any
args: tuple[Any, ...]
kwargs: dict[str, Any]
transforms: tuple[Any, ...]
random_before: tuple[mx.array, ...] | None
compile_allowed: bool
@dataclass(frozen=True)
class _LossPlan:
model: Any
cache_key: tuple[Any, ...]
dynamic_inputs: tuple[mx.array, ...]
template: Any
transforms: tuple[Any, ...]
rebuild: Callable[..., mx.array]
random_before: tuple[mx.array, ...] | None
compile_allowed: bool
class _TrainingPlan:
def __init__(self, model, optimizer, objective, compile=True):
def __init__(
self,
model: Any,
optimizer: Any,
objective: Callable[..., mx.array],
compile: bool = True,
) -> None:
self.model = model
self.optimizer = optimizer
self._value_and_grad = nn.value_and_grad(model, objective)
self._compile_enabled = compile
self._compiled_backward = None
self._compiled_fused = None
self._backward_signatures = set()
self._fused_signatures = set()
self._compiled_backward: Callable[..., Any] | None = None
self._compiled_fused: Callable[..., Any] | None = None
self._backward_signatures: set[Any] = set()
self._fused_signatures: set[Any] = set()
if compile:
backward_state = [model.state, mx.random.state]
fused_state = [model.state, optimizer.state, mx.random.state]
self._compiled_backward = mx.compile(
self._value_and_grad,
inputs=backward_state,
outputs=backward_state,
self._value_and_grad, inputs=backward_state, outputs=backward_state
)
self._compiled_fused = mx.compile(
self._fused,
inputs=fused_state,
outputs=fused_state,
self._fused, inputs=fused_state, outputs=fused_state
)
def _fused(self, *inputs):
def _fused(self, *inputs: mx.array) -> mx.array:
loss, gradients = self._value_and_grad(*inputs)
self.optimizer.update(self.model, gradients)
return loss
def backward(self, *inputs):
def backward(self, *inputs: mx.array) -> Any:
if self._compile_enabled:
assert self._compiled_backward is not None
return self._compiled_backward(*inputs)
return self._value_and_grad(*inputs)
def fused(self, *inputs):
def fused(self, *inputs: mx.array) -> mx.array:
if self._compile_enabled:
assert self._compiled_fused is not None
return self._compiled_fused(*inputs)
return self._fused(*inputs)
def disable_compile(self):
def disable_compile(self) -> None:
self._compile_enabled = False
def _partition_arrays(value, arrays):
def _partition_arrays(value: Any, arrays: list[mx.array]) -> Any:
if isinstance(value, mx.array):
slot = _ArraySlot(len(arrays))
arrays.append(value)
@@ -72,7 +105,7 @@ def _partition_arrays(value, arrays):
return value
def _restore_arrays(template, arrays):
def _restore_arrays(template: Any, arrays: tuple[mx.array, ...]) -> Any:
if isinstance(template, _ArraySlot):
return arrays[template.index]
if isinstance(template, tuple):
@@ -84,7 +117,7 @@ def _restore_arrays(template, arrays):
return template
def _template_key(template):
def _template_key(template: Any) -> tuple[Any, ...]:
if isinstance(template, _ArraySlot):
return ("array",)
if isinstance(template, tuple):
@@ -102,7 +135,7 @@ def _template_key(template):
return ("static", type(template).__qualname__, repr(template))
def _copy_tree(value):
def _copy_tree(value: Any) -> Any:
if isinstance(value, dict):
return {key: _copy_tree(item) for key, item in value.items()}
if isinstance(value, list):
@@ -112,211 +145,198 @@ def _copy_tree(value):
return value
def _snapshot_random_state():
state = [key + mx.array(0, dtype=key.dtype) for key in mx.random.state]
def _snapshot_random_state() -> list[mx.array]:
random_state = cast(Any, mx.random.state)
state = [key + mx.array(0, dtype=key.dtype) for key in random_state]
mx.eval(state)
return state
def _restore_random_state(state):
for current_key, saved_key in zip(mx.random.state, state):
def _restore_random_state(state: Any) -> None:
for current_key, saved_key in zip(cast(Any, mx.random.state), state):
current_key[...] = saved_key
mx.eval(mx.random.state)
def _compile_failure(error):
def _compile_failure(error: Exception) -> bool:
message = str(error)
return "Attempting to eval an array" in message and (
"function transformations" in message or "without a primitive" in message
)
def begin_forward(model):
global _forward_depth
if (
_suspended
or _active_optimizer is None
or model is not _active_optimizer._model
):
def begin_forward(model: Any) -> bool:
optimizer = _active_optimizer.get()
if _suspended.get() or optimizer is None or model is not optimizer._model:
return False
if _forward_depth:
_active_optimizer._recursive_forward = True
_forward_depth += 1
depth = _forward_depth.get()
if depth:
optimizer._recursive_forward = True
_forward_depth.set(depth + 1)
return True
def note_eager_evaluation():
if _forward_depth and _active_optimizer is not None and not _suspended:
_active_optimizer._eager_forward = True
def note_eager_evaluation() -> None:
optimizer = _active_optimizer.get()
if _forward_depth.get() and optimizer is not None and not _suspended.get():
optimizer._eager_forward = True
def end_forward(model, args, kwargs, output):
global _forward_depth
_forward_depth -= 1
if _forward_depth != 0 or not isinstance(output, mx.array):
def end_forward(
model: Any, args: tuple[Any, ...], kwargs: dict[str, Any], output: Any
) -> None:
optimizer = _active_optimizer.get()
depth = _forward_depth.get() - 1
_forward_depth.set(depth)
if depth != 0 or optimizer is None or not isinstance(output, Tensor):
return
random_before = _active_optimizer._random_before_forward
random_before = optimizer._random_before_forward
stochastic = any(
current is not previous
for current, previous in zip(mx.random.state, random_before)
for current, previous in zip(cast(Any, mx.random.state), random_before)
)
descriptor = (
output._lineage = _Lineage(
model,
args,
kwargs,
(),
random_before if stochastic else None,
not (
_active_optimizer._recursive_forward
or _active_optimizer._eager_forward
),
not (optimizer._recursive_forward or optimizer._eager_forward),
)
_lineage[id(output)] = (output, descriptor)
def abort_forward():
global _forward_depth
_forward_depth -= 1
def abort_forward() -> None:
depth = max(0, _forward_depth.get() - 1)
_forward_depth.set(depth)
if depth == 0:
optimizer = _active_optimizer.get()
if optimizer is not None:
optimizer._pending_update = None
_active_optimizer.set(None)
def activate(optimizer):
global _active_optimizer, _forward_depth
_active_optimizer = optimizer
_forward_depth = 0
_loss_plans.clear()
_lineage.clear()
_replayed_losses.clear()
optimizer._random_before_forward = tuple(mx.random.state)
def activate(optimizer: Any) -> None:
previous = _active_optimizer.get()
if previous is not None and previous is not optimizer:
previous._pending_update = None
_active_optimizer.set(optimizer)
_forward_depth.set(0)
optimizer._random_before_forward = tuple(cast(Any, mx.random.state))
optimizer._recursive_forward = False
optimizer._eager_forward = False
def propagate(source, result, operation, signature, operands=()):
if _suspended:
def propagate(
source: Tensor,
result: Tensor,
operation: Callable[..., mx.array],
signature: tuple[Any, ...],
operands: tuple[mx.array, ...] = (),
) -> Tensor:
if _suspended.get() or source._lineage is None:
return result
entry = _lineage.get(id(source))
if entry is None or entry[0] is not source:
return result
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
_lineage[id(result)] = (
result,
(
model,
args,
kwargs,
transforms + ((operation, signature, operands),),
random_before,
compile_allowed,
),
lineage: _Lineage = source._lineage
result._lineage = _Lineage(
lineage.model,
lineage.args,
lineage.kwargs,
lineage.transforms + ((operation, signature, operands),),
lineage.random_before,
lineage.compile_allowed,
)
return result
def _plan_for(
optimizer, model, cache_key, template, transforms, rebuild, compile_allowed
):
plan = optimizer._training_plans.get(cache_key)
def _plan_for(optimizer: Any, loss_plan: _LossPlan) -> _TrainingPlan:
plan = optimizer._training_plans.get(loss_plan.cache_key)
if plan is not None:
return plan
def objective(*current_inputs):
current_args, current_kwargs, current_transform_operands, current_operands = (
_restore_arrays(template, current_inputs)
def objective(*current_inputs: mx.array) -> mx.array:
current_args, current_kwargs, transform_operands, current_operands = (
_restore_arrays(loss_plan.template, current_inputs)
)
output = model(*current_args, **current_kwargs)
for (operation, _, _), operation_operands in zip(
transforms, current_transform_operands
output = loss_plan.model(*current_args, **current_kwargs)
for (operation, _, _), operands in zip(
loss_plan.transforms, transform_operands
):
output = operation(output, *operation_operands)
return rebuild(output, *current_operands)
output = operation(output, *operands)
return loss_plan.rebuild(output, *current_operands)
plan = _TrainingPlan(
model, optimizer._optimizer, objective, compile=compile_allowed
loss_plan.model,
optimizer._optimizer,
objective,
compile=loss_plan.compile_allowed,
)
optimizer._training_plans[cache_key] = plan
optimizer._training_plans[loss_plan.cache_key] = plan
return plan
def register_loss(loss, loss_input, rebuild, operands, signature):
if _suspended or _active_optimizer is None:
def register_loss(
loss: Tensor,
loss_input: Tensor,
rebuild: Callable[..., mx.array],
operands: tuple[mx.array, ...],
signature: tuple[Any, ...],
) -> Tensor:
optimizer = _active_optimizer.get()
if _suspended.get() or optimizer is None or loss_input._lineage is None:
return loss
entry = _lineage.get(id(loss_input))
if entry is None or entry[0] is not loss_input:
return loss
model, args, kwargs, transforms, random_before, compile_allowed = entry[1]
lineage: _Lineage = loss_input._lineage
context = (
args,
kwargs,
tuple(transform_operands for _, _, transform_operands in transforms),
lineage.args,
lineage.kwargs,
tuple(cast(Any, item)[2] for item in lineage.transforms),
operands,
)
dynamic_inputs = []
dynamic_inputs: list[mx.array] = []
template = _partition_arrays(context, dynamic_inputs)
cache_key = (
getattr(model, "training", None),
getattr(lineage.model, "training", None),
_template_key(template),
tuple(transform_signature for _, transform_signature, _ in transforms),
tuple(item[1] for item in lineage.transforms),
signature,
compile_allowed,
lineage.compile_allowed,
)
_loss_plans[id(loss)] = (
loss,
model,
loss._loss_plan = _LossPlan(
lineage.model,
cache_key,
tuple(dynamic_inputs),
template,
transforms,
lineage.transforms,
rebuild,
random_before,
compile_allowed,
lineage.random_before,
lineage.compile_allowed,
)
return loss
def backward(loss):
if _active_optimizer is None:
def backward(loss: Tensor) -> None:
optimizer = _active_optimizer.get()
if optimizer is None:
raise RuntimeError("optimizer.zero_grad() must be called before loss.backward()")
entry = _loss_plans.get(id(loss))
if entry is None or entry[0] is not loss:
loss_plan = loss._loss_plan
if loss_plan is None:
_clear(optimizer)
raise RuntimeError(
"this MLX loss cannot use backward compatibility; compute it with a supported torchmlx loss function"
"this MLX loss cannot use backward compatibility; compute one supported loss from one model forward"
)
(
_,
model,
cache_key,
dynamic_inputs,
template,
transforms,
rebuild,
random_before,
compile_allowed,
) = entry
plan = _plan_for(
_active_optimizer,
model,
cache_key,
template,
transforms,
rebuild,
compile_allowed,
)
_active_optimizer._pending_update = (
try:
plan = _plan_for(optimizer, loss_plan)
except Exception:
_clear(optimizer)
raise
optimizer._pending_update = (
"staged",
loss,
plan,
dynamic_inputs,
random_before,
loss_plan.dynamic_inputs,
loss_plan.random_before,
)
def _restore_execution(model, optimizer, parameters, optimizer_state):
model.update(parameters)
if optimizer_state is not None:
_restore_tree(optimizer.state, optimizer_state)
def _restore_tree(current, saved):
def _restore_tree(current: Any, saved: Any) -> Any:
if isinstance(current, dict) and isinstance(saved, dict):
for key in tuple(current):
if key not in saved:
@@ -330,16 +350,26 @@ def _restore_tree(current, saved):
current[key] = value
return current
if isinstance(current, list) and isinstance(saved, list):
restored = [
_restore_tree(old, value) for old, value in zip(current, saved)
]
restored = [_restore_tree(old, value) for old, value in zip(current, saved)]
current[:] = restored + saved[len(restored) :]
return current
return saved
def _execute(plan, inputs, fused, random_before, advance_random=False):
global _suspended
def _restore_execution(
model: Any, optimizer: Any, parameters: Any, optimizer_state: Any
) -> None:
model.update(parameters)
if optimizer_state is not None:
_restore_tree(optimizer.state, optimizer_state)
def _execute(
plan: _TrainingPlan,
inputs: tuple[mx.array, ...],
fused: bool,
random_before: Any,
) -> Any:
model = plan.model
optimizer = plan.optimizer
signatures = plan._fused_signatures if fused else plan._backward_signatures
@@ -358,22 +388,19 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
else getattr(optimizer, "_torchmlx_state", None)
)
random_after = _snapshot_random_state() if random_before is not None else None
random_rollback = tuple(mx.random.state) if advance_random else None
if random_before is not None:
_restore_random_state(random_before)
_suspended = True
token = _suspended.set(True)
try:
try:
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
current_parameters = model.parameters()
current_parameters = model._native_parameters()
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
mx.eval(
result,
current_parameters,
current_optimizer_state,
mx.random.state
if random_before is not None or advance_random
else [],
mx.random.state if random_before is not None else [],
)
signatures.add(signature)
if fused:
@@ -385,19 +412,15 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
_restore_execution(model, optimizer, parameters, optimizer_state)
if random_before is not None:
_restore_random_state(random_before)
elif random_rollback is not None:
_restore_random_state(random_rollback)
plan.disable_compile()
result = plan.fused(*inputs) if fused else plan.backward(*inputs)
current_parameters = model.parameters()
current_parameters = model._native_parameters()
current_optimizer_state = _copy_tree(optimizer.state) if fused else []
mx.eval(
result,
current_parameters,
current_optimizer_state,
mx.random.state
if random_before is not None or advance_random
else [],
mx.random.state if random_before is not None else [],
)
if fused:
optimizer._torchmlx_parameters = current_parameters
@@ -405,105 +428,65 @@ def _execute(plan, inputs, fused, random_before, advance_random=False):
except Exception:
if parameters is not None:
_restore_execution(model, optimizer, parameters, optimizer_state)
if random_rollback is not None:
_restore_random_state(random_rollback)
raise
finally:
_suspended = False
_suspended.reset(token)
if random_after is not None:
_restore_random_state(random_after)
return result
def _materialize(optimizer):
pending = optimizer._pending_update
_, loss, plan, dynamic_inputs, random_before = pending
def _materialize(optimizer: Any) -> mx.array:
_, loss, plan, dynamic_inputs, random_before = optimizer._pending_update
replayed_loss, gradients = _execute(
plan, dynamic_inputs, fused=False, random_before=random_before
)
_replayed_losses[id(loss)] = (loss, replayed_loss)
loss._replayed_value = replayed_loss
optimizer._pending_update = ("materialized", loss, plan.model, gradients)
return replayed_loss
def replayed_value(value):
if _active_optimizer is not None:
pending = _active_optimizer._pending_update
def replayed_value(value: Tensor) -> mx.array:
optimizer = _active_optimizer.get()
if optimizer is not None:
pending = optimizer._pending_update
if pending is not None and pending[0] == "staged" and pending[1] is value:
return _materialize(_active_optimizer)
entry = _replayed_losses.get(id(value))
if entry is None or entry[0] is not value:
return value
return entry[1]
return _materialize(optimizer)
return value._replayed_value if value._replayed_value is not None else value
def _finish_step(loss, replayed_loss):
global _active_optimizer, _forward_depth
_replayed_losses[id(loss)] = (loss, replayed_loss)
_active_optimizer._pending_update = None
_active_optimizer = None
_forward_depth = 0
_loss_plans.clear()
_lineage.clear()
def _clear(optimizer: Any) -> None:
optimizer._pending_update = None
_active_optimizer.set(None)
_forward_depth.set(0)
def step(optimizer):
def step(optimizer: Any) -> None:
pending = optimizer._pending_update
if pending is None:
_clear(optimizer)
raise RuntimeError("loss.backward() must be called before optimizer.step()")
if pending[0] == "staged":
_, loss, plan, dynamic_inputs, random_before = pending
replayed_loss = _execute(
plan, dynamic_inputs, fused=True, random_before=random_before
)
else:
_, loss, model, gradients = pending
parameters = _copy_tree(model.trainable_parameters())
optimizer_state = _copy_tree(optimizer._optimizer.state)
try:
optimizer._compiled_step(gradients)
current_parameters = model.parameters()
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
mx.eval(current_parameters, current_optimizer_state)
except Exception:
_restore_execution(
model, optimizer._optimizer, parameters, optimizer_state
try:
if pending[0] == "staged":
_, loss, plan, dynamic_inputs, random_before = pending
replayed_loss = _execute(
plan, dynamic_inputs, fused=True, random_before=random_before
)
raise
optimizer._optimizer._torchmlx_parameters = current_parameters
optimizer._optimizer._torchmlx_state = current_optimizer_state
replayed_loss = _replayed_losses[id(loss)][1]
_finish_step(loss, replayed_loss)
def trainer_step(trainer, x, y):
context = (x, y)
dynamic_inputs = []
template = _partition_arrays(context, dynamic_inputs)
cache_key = (
getattr(trainer.model, "training", None),
_template_key(template),
(),
("trainer", id(trainer.loss_fn), trainer.compile),
)
plan = trainer.optimizer._training_plans.get(cache_key)
if plan is None:
def objective(*current_inputs):
current_x, current_y = _restore_arrays(template, current_inputs)
return trainer.loss_fn(trainer.model(current_x), current_y)
plan = _TrainingPlan(
trainer.model,
trainer.optimizer._optimizer,
objective,
compile=trainer.compile,
)
trainer.optimizer._training_plans[cache_key] = plan
return _execute(
plan,
tuple(dynamic_inputs),
fused=True,
random_before=None,
advance_random=True,
)
else:
_, loss, model, gradients = pending
parameters = _copy_tree(model.trainable_parameters())
optimizer_state = _copy_tree(optimizer._optimizer.state)
try:
optimizer._compiled_step(gradients)
current_parameters = model._native_parameters()
current_optimizer_state = _copy_tree(optimizer._optimizer.state)
mx.eval(current_parameters, current_optimizer_state)
except Exception:
_restore_execution(model, optimizer._optimizer, parameters, optimizer_state)
raise
optimizer._optimizer._torchmlx_parameters = current_parameters
optimizer._optimizer._torchmlx_state = current_optimizer_state
replayed_loss = loss._replayed_value
loss._replayed_value = replayed_loss
finally:
_clear(optimizer)
+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}")
+311 -253
View File
@@ -1,275 +1,333 @@
from __future__ import annotations
from typing import Any, Callable, cast
import mlx.core as mx
_transpose = mx.array.transpose
_reshape = mx.array.reshape
_squeeze = mx.array.squeeze
_mean = mx.array.mean
_var = mx.array.var
_any = mx.array.any
_getitem = mx.array.__getitem__
_setitem = mx.array.__setitem__
_item = mx.array.item
_native_getitem = mx.array.__getitem__
_native_setitem = mx.array.__setitem__
_native_item = mx.array.item
_native_reshape = mx.array.reshape
_native_squeeze = mx.array.squeeze
_native_transpose = mx.array.transpose
def _torch_transpose(self, dim0=None, dim1=None):
if dim1 is None:
if dim0 is None:
axes = list(reversed(range(self.ndim)))
def wrap(value: Any) -> Any:
if isinstance(value, Tensor):
return value
if isinstance(value, mx.array):
return Tensor(value)
if isinstance(value, tuple):
return tuple(wrap(item) for item in value)
if isinstance(value, list):
return [wrap(item) for item in value]
if isinstance(value, dict):
return {key: wrap(item) for key, item in value.items()}
return value
class Tensor(mx.array):
_lineage: Any = None
_loss_plan: Any = None
_replayed_value: mx.array | None = None
def _result(
self,
value: Any,
operation: Callable[..., mx.array] | None = None,
signature: tuple[Any, ...] = (),
operands: tuple[mx.array, ...] = (),
) -> Any:
result = wrap(value)
if operation is not None and isinstance(result, Tensor):
from ._autograd import propagate
propagate(self, result, operation, signature, operands)
return result
@property
def device(self) -> Any:
from ._mlx_backend import device
return device("mps")
@property
def grad(self) -> Any:
raise RuntimeError("parameter gradients are not exposed by the MLX backend")
def register_hook(self, hook: Any) -> Any:
raise RuntimeError("gradient hooks are not supported by the MLX backend")
def transpose(self, dim0: Any = None, dim1: int | None = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
if dim1 is None:
axes = list(reversed(range(self.ndim))) if dim0 is None else dim0
else:
axes = dim0
else:
axes = list(range(self.ndim))
axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
result = _transpose(self, axes)
from ._autograd import propagate
return propagate(
self,
result,
lambda value: _transpose(value, axes),
("transpose", tuple(axes)),
)
def _view(self, *shape):
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
shape = shape[0]
return _torch_reshape(self, shape)
def _torch_reshape(self, *shape):
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
shape = shape[0]
result = _reshape(self, shape)
from ._autograd import propagate
return propagate(
self,
result,
lambda value: _reshape(value, shape),
("reshape", tuple(shape)),
)
def _unsqueeze(self, dim):
result = mx.expand_dims(self, axis=dim)
from ._autograd import propagate
return propagate(
self,
result,
lambda value: mx.expand_dims(value, axis=dim),
("unsqueeze", dim),
)
def _torch_squeeze(self, dim=None):
result = _squeeze(self, axis=dim)
from ._autograd import propagate
return propagate(
self,
result,
lambda value: _squeeze(value, axis=dim),
("squeeze", dim),
)
def _flatten(self, start_dim=0, end_dim=-1):
if end_dim < 0:
end_dim += self.ndim
flattened = 1
for dimension in self.shape[start_dim : end_dim + 1]:
flattened *= dimension
shape = self.shape[:start_dim] + (flattened,) + self.shape[end_dim + 1 :]
return _torch_reshape(self, shape)
def _float(self):
return self.astype(mx.float32)
def _bool(self):
return self.astype(mx.bool_)
def _pow(self, exponent):
return mx.power(self, exponent)
def _torch_mean(self, dim=None, keepdim=False, dtype=None):
value = self.astype(dtype) if dtype is not None else self
return _mean(value, axis=dim, keepdims=keepdim)
def _torch_var(
self,
dim=None,
unbiased=True,
keepdim=False,
*,
correction=None,
):
ddof = int(unbiased) if correction is None else correction
return _var(self, axis=dim, keepdims=keepdim, ddof=ddof)
def _size(self, dim=None):
return self.shape if dim is None else self.shape[dim]
def _chunk(self, chunks, dim=0):
if chunks <= 0:
raise ValueError("chunks must be greater than 0")
length = self.shape[dim]
if length == 0:
return tuple(mx.split(self, chunks, axis=dim))
chunk_size = (length + chunks - 1) // chunks
indices = list(range(chunk_size, length, chunk_size))
return tuple(mx.split(self, indices, axis=dim))
def _to(self, *args, dtype=None, device=None, **kwargs):
if kwargs:
name = next(iter(kwargs))
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
for value in args:
if isinstance(value, mx.Dtype):
dtype = value
elif isinstance(value, mx.array):
dtype = value.dtype
else:
device = value
if device is not None and str(device) != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
return self.astype(dtype) if dtype is not None and dtype != self.dtype else self
def _repeat_interleave(self, repeats, dim=None):
return mx.repeat(self, repeats, axis=dim)
def _masked_fill(self, mask, value):
return mx.where(mask, mx.array(value, dtype=self.dtype), self)
def _expand(self, *sizes):
if len(sizes) == 1 and isinstance(sizes[0], (tuple, list)):
sizes = tuple(sizes[0])
if len(sizes) < self.ndim:
raise ValueError("expanded size must have at least as many dimensions as the tensor")
source = (1,) * (len(sizes) - self.ndim) + self.shape
target = tuple(current if requested == -1 else requested for requested, current in zip(sizes, source))
return mx.broadcast_to(self.reshape(source), target)
def _torch_any(self, dim=None, keepdim=False):
return _any(self, axis=dim, keepdims=keepdim)
def _contiguous(self, memory_format=None):
return self
def _requires_grad(self, requires_grad=True):
if requires_grad:
raise RuntimeError(
"requires_grad_ is not supported by the MLX backend; use torchmlx.Trainer"
axes = list(range(self.ndim))
axes[dim0], axes[dim1] = axes[dim1], axes[dim0]
return self._result(
_native_transpose(self, axes),
lambda value: _native_transpose(value, axes),
("transpose", tuple(axes)),
)
return self
def view(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
return self.reshape(*shape)
def _backward(self, *args, **kwargs):
if args or kwargs:
raise TypeError("MLX backward compatibility does not accept arguments")
from ._autograd import backward
def reshape(self, *shape: Any) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
if len(shape) == 1 and isinstance(shape[0], (tuple, list)):
shape = tuple(shape[0])
return self._result(
_native_reshape(self, shape),
lambda value: _native_reshape(value, shape),
("reshape", tuple(shape)),
)
backward(self)
def unsqueeze(self, dim: int) -> Tensor:
return self._result(
mx.expand_dims(self, axis=dim),
lambda value: mx.expand_dims(value, axis=dim),
("unsqueeze", dim),
)
def squeeze(self, dim: int | None = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
return self._result(
_native_squeeze(self, axis=dim),
lambda value: _native_squeeze(value, axis=dim),
("squeeze", dim),
)
def _torch_item(self):
from ._autograd import note_eager_evaluation, replayed_value
def flatten(self, start_dim: int = 0, end_dim: int = -1) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
if end_dim < 0:
end_dim += self.ndim
flattened = 1
for dimension in self.shape[start_dim : end_dim + 1]:
flattened *= dimension
shape = self.shape[:start_dim] + (flattened,) + self.shape[end_dim + 1 :]
return self.reshape(shape)
note_eager_evaluation()
return _item(replayed_value(self))
def float(self) -> Tensor:
return self.astype(mx.float32)
def bool(self) -> Tensor:
return self.astype(mx.bool_)
def _mask_indices(mask):
flat = mask.reshape(-1)
count = int(mx.sum(flat).item())
order = mx.argsort(flat.astype(mx.int32))
if count == 0:
return order[:0].astype(mx.int64)
return order[-count:].astype(mx.int64)
def pow(self, exponent: Any) -> Tensor:
return self._binary(mx.power, exponent, "pow")
def mean( # pyright: ignore[reportIncompatibleMethodOverride]
self, dim: Any = None, keepdim: bool = False, dtype: Any = None
) -> Tensor:
value = mx.astype(self, dtype) if dtype is not None else self
return wrap(mx.mean(value, axis=dim, keepdims=keepdim))
def _torch_getitem(self, key):
if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key)
if key.shape == self.shape:
result = _getitem(self.reshape(-1), indices)
operation = lambda value, current_indices: _getitem(
_reshape(value, (-1,)), current_indices
def var( # pyright: ignore[reportIncompatibleMethodOverride]
self,
dim: Any = None,
unbiased: bool = True,
keepdim: bool = False,
*,
correction: int | None = None,
) -> Tensor:
ddof = int(unbiased) if correction is None else correction
return wrap(mx.var(self, axis=dim, keepdims=keepdim, ddof=ddof))
def size(self, dim: int | None = None) -> Any: # pyright: ignore[reportIncompatibleMethodOverride]
return self.shape if dim is None else self.shape[dim]
def chunk(self, chunks: int, dim: int = 0) -> tuple[Tensor, ...]:
if chunks <= 0:
raise ValueError("chunks must be greater than 0")
length = self.shape[dim]
if length == 0:
return wrap(tuple(mx.split(self, chunks, axis=dim)))
chunk_size = (length + chunks - 1) // chunks
indices = list(range(chunk_size, length, chunk_size))
return wrap(tuple(mx.split(self, indices, axis=dim)))
def to(
self,
*args: Any,
dtype: Any = None,
device: Any = None,
**kwargs: Any,
) -> Tensor:
if kwargs:
name = next(iter(kwargs))
raise TypeError(f"to() got an unexpected keyword argument {name!r}")
for value in args:
if isinstance(value, mx.Dtype):
dtype = value
elif isinstance(value, mx.array):
dtype = value.dtype
else:
device = value
if device is not None and str(device) != "mps":
raise ValueError("the MLX backend only accepts device='mps'")
return self.astype(dtype) if dtype is not None and dtype != self.dtype else self
def repeat_interleave(self, repeats: int, dim: int | None = None) -> Tensor:
return wrap(mx.repeat(self, repeats, axis=dim))
def masked_fill(self, mask: mx.array, value: Any) -> Tensor:
return wrap(mx.where(mask, mx.array(value, dtype=self.dtype), self))
def expand(self, *sizes: Any) -> Tensor:
if len(sizes) == 1 and isinstance(sizes[0], (tuple, list)):
sizes = tuple(sizes[0])
if len(sizes) < self.ndim:
raise ValueError("expanded size must have at least as many dimensions as the tensor")
source = (1,) * (len(sizes) - self.ndim) + self.shape
target = tuple(
current if requested == -1 else requested
for requested, current in zip(sizes, source)
)
return wrap(mx.broadcast_to(_native_reshape(self, source), target))
def any(self, dim: Any = None, keepdim: bool = False) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
return wrap(mx.any(self, axis=dim, keepdims=keepdim))
def contiguous(self, memory_format: Any = None) -> Tensor:
return self
def requires_grad_(self, requires_grad: bool = True) -> Tensor:
if requires_grad:
raise RuntimeError("requires_grad_ is not supported by the MLX backend")
return self
def backward(self, *args: Any, **kwargs: Any) -> None:
if args or kwargs:
if kwargs.get("create_graph") or kwargs.get("retain_graph"):
raise RuntimeError("higher-order gradients are not supported by the MLX backend")
raise TypeError("MLX backward compatibility does not accept arguments")
from ._autograd import backward
backward(self)
def item(self) -> Any:
from ._autograd import note_eager_evaluation, replayed_value
note_eager_evaluation()
return _native_item(replayed_value(self))
def astype(self, dtype: Any, *, stream: Any = None) -> Tensor: # pyright: ignore[reportIncompatibleMethodOverride]
kwargs = {} if stream is None else {"stream": stream}
return wrap(mx.astype(self, dtype, **kwargs))
def __getitem__(self, key: Any) -> Tensor:
if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key)
if key.shape == self.shape:
result = _native_getitem(_native_reshape(self, (-1,)), indices)
operation = lambda value, current: _native_getitem(
_native_reshape(value, (-1,)), current
)
signature = ("getitem_bool", "flat")
else:
result = _native_getitem(self, indices)
operation = lambda value, current: _native_getitem(value, current)
signature = ("getitem_bool", "first_axis")
return self._result(result, operation, signature, (indices,))
if isinstance(key, mx.array):
return self._result(
_native_getitem(self, key),
lambda value, current: _native_getitem(value, current),
("getitem_array",),
(key,),
)
signature = ("getitem_bool", "flat")
else:
result = _getitem(self, indices)
operation = lambda value, current_indices: _getitem(
value, current_indices
)
signature = ("getitem_bool", "first_axis")
operands = (indices,)
elif isinstance(key, mx.array):
result = _getitem(self, key)
operation = lambda value, current_key: _getitem(value, current_key)
signature = ("getitem_array",)
operands = (key,)
else:
result = _getitem(self, key)
operation = lambda value: _getitem(value, key)
signature = ("getitem", type(key).__qualname__, repr(key))
operands = ()
from ._autograd import propagate
return self._result(
_native_getitem(self, key),
lambda value: _native_getitem(value, key),
("getitem", type(key).__qualname__, repr(key)),
)
return propagate(self, result, operation, signature, operands)
def _torch_setitem(self, key, value):
if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key)
if key.shape == self.shape:
flat = self.reshape(-1)
_setitem(flat, indices, value)
def __setitem__(self, key: Any, value: Any) -> None:
if isinstance(key, mx.array) and key.dtype == mx.bool_:
indices = _mask_indices(key)
if key.shape == self.shape:
_native_setitem(_native_reshape(self, (-1,)), indices, value)
return
_native_setitem(self, indices, value)
return
_setitem(self, indices, value)
return
_setitem(self, key, value)
_native_setitem(self, key, value)
def _binary(self, function: Callable[..., mx.array], other: Any, name: str) -> Tensor:
if isinstance(other, Tensor) and other._lineage is not None:
return wrap(function(self, other))
operands = (other,) if isinstance(other, mx.array) else ()
operation = (
(lambda value, current: function(value, current))
if operands
else (lambda value: function(value, other))
)
signature = (name, "array" if operands else repr(other))
return self._result(function(self, other), operation, signature, operands)
def _reverse(self, function: Callable[..., mx.array], other: Any) -> Tensor:
return wrap(function(other, self))
def __add__(self, other: Any) -> Tensor:
return self._binary(mx.add, other, "add")
def __radd__(self, other: Any) -> Tensor:
return self._reverse(mx.add, other)
def __sub__(self, other: Any) -> Tensor:
return self._binary(mx.subtract, other, "sub")
def __rsub__(self, other: Any) -> Tensor:
return self._reverse(mx.subtract, other)
def __mul__(self, other: Any) -> Tensor:
return self._binary(mx.multiply, other, "mul")
def __rmul__(self, other: Any) -> Tensor:
return self._reverse(mx.multiply, other)
def __truediv__(self, other: Any) -> Tensor:
return self._binary(mx.divide, other, "div")
def __rtruediv__(self, other: Any) -> Tensor:
return self._reverse(mx.divide, other)
def __pow__(self, other: Any) -> Tensor:
return self.pow(other)
def __rpow__(self, other: Any) -> Tensor:
return self._reverse(mx.power, other)
def __matmul__(self, other: Any) -> Tensor:
return self._binary(mx.matmul, other, "matmul")
def __rmatmul__(self, other: Any) -> Tensor:
return self._reverse(mx.matmul, other)
def __neg__(self) -> Tensor:
return wrap(mx.negative(self))
def __eq__(self, other: Any) -> Tensor:
return wrap(mx.equal(self, other))
def __ne__(self, other: Any) -> Tensor:
return wrap(mx.not_equal(self, other))
def __lt__(self, other: Any) -> Tensor:
return wrap(mx.less(self, other))
def __le__(self, other: Any) -> Tensor:
return wrap(mx.less_equal(self, other))
def __gt__(self, other: Any) -> Tensor:
return wrap(mx.greater(self, other))
def __ge__(self, other: Any) -> Tensor:
return wrap(mx.greater_equal(self, other))
def install(device_type):
mx.array.transpose = _torch_transpose
mx.array.reshape = _torch_reshape
mx.array.view = _view
mx.array.unsqueeze = _unsqueeze
mx.array.squeeze = _torch_squeeze
mx.array.flatten = _flatten
mx.array.float = _float
mx.array.bool = _bool
mx.array.pow = _pow
mx.array.mean = _torch_mean
mx.array.var = _torch_var
mx.array.size = _size
mx.array.chunk = _chunk
mx.array.to = _to
mx.array.repeat_interleave = _repeat_interleave
mx.array.masked_fill = _masked_fill
mx.array.expand = _expand
mx.array.any = _torch_any
mx.array.contiguous = _contiguous
mx.array.requires_grad_ = _requires_grad
mx.array.backward = _backward
mx.array.item = _torch_item
mx.array.device = property(lambda self: device_type("mps"))
mx.array.__getitem__ = _torch_getitem
mx.array.__setitem__ = _torch_setitem
def _mask_indices(mask: mx.array) -> mx.array:
flat = _native_reshape(mask, (-1,))
count = int(cast(Any, _native_item(mx.sum(flat))))
order = mx.argsort(mx.astype(flat, mx.int32))
if count == 0:
return mx.astype(_native_getitem(order, slice(0, 0)), mx.int64)
return mx.astype(_native_getitem(order, slice(-count, None)), mx.int64)
+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
@@ -19,26 +23,62 @@ if BACKEND == "torch":
else:
from collections.abc import Iterable, MutableSequence
from contextvars import ContextVar
import mlx.core as mx
import mlx.nn as _native
from mlx.utils import tree_flatten
class _ParameterTree(dict):
from torchmlx._mlx_tensor import Tensor, wrap
_module_depth = ContextVar("torchmlx_module_depth", default=0)
class _ParameterIterator:
def __init__(self, values, model):
super().__init__(values)
self.model = model
self._iterator = iter(value for _, value in tree_flatten(values))
def __iter__(self):
return self
def __next__(self):
return wrap(next(self._iterator))
class Module(_native.Module):
def __setattr__(self, name, value):
super().__setattr__(name, mx.array(value) if isinstance(value, Tensor) else value)
def __getattribute__(self, name):
value = super().__getattribute__(name)
return (
wrap(value)
if isinstance(value, mx.array) and _module_depth.get() == 0
else value
)
def __getattr__(self, name):
value = super().__getattr__(name)
return (
wrap(value)
if isinstance(value, mx.array) and _module_depth.get() == 0
else value
)
def __call__(self, *args, **kwargs):
from torchmlx._autograd import abort_forward, begin_forward, end_forward
tracking = begin_forward(self)
args = wrap(args)
kwargs = wrap(kwargs)
token = _module_depth.set(_module_depth.get() + 1)
try:
output = self.forward(*args, **kwargs)
output = wrap(self.forward(*args, **kwargs))
except Exception:
if tracking:
abort_forward()
raise
finally:
_module_depth.reset(token)
if tracking:
end_forward(self, args, kwargs, output)
return output
@@ -49,7 +89,10 @@ else:
)
def parameters(self):
return _ParameterTree(super().parameters(), self)
return _ParameterIterator(self._native_parameters(), self)
def _native_parameters(self):
return super().parameters()
def to(self, *args, **kwargs):
dtype = kwargs.pop("dtype", None)
@@ -80,10 +123,10 @@ else:
def Parameter(data=None, requires_grad=True):
if data is None:
return mx.array([])
return wrap(mx.array([]))
if not requires_grad:
unsupported("torchmlx.nn.Parameter with requires_grad=False")
return data
return wrap(data)
class Linear(Module, _native.Linear):
def __init__(self, in_features, out_features, bias=True, device=None, dtype=None):
@@ -156,7 +199,7 @@ else:
raise ValueError("the MLX backend only accepts device='mps'")
_native.LayerNorm.__init__(
self,
normalized_shape,
cast(int, normalized_shape),
eps=eps,
affine=elementwise_affine,
bias=bias,
+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: ...
+10 -5
View File
@@ -1,3 +1,5 @@
# pyright: reportAssignmentType=false, reportRedeclaration=false
import math
from torchmlx._backend import BACKEND, unsupported
@@ -16,6 +18,7 @@ if BACKEND == "torch":
else:
import mlx.core as mx
from torchmlx._mlx_tensor import wrap
def cross_entropy(
input,
@@ -55,7 +58,7 @@ else:
label_smoothing,
)
return register_loss(
loss,
wrap(loss),
original_input,
rebuild,
operands,
@@ -125,18 +128,20 @@ else:
if scale is None:
scale = 1 / math.sqrt(query.shape[-1])
mask = "causal" if is_causal else attn_mask
return mx.fast.scaled_dot_product_attention(
query, key, value, scale=scale, mask=mask
return wrap(
mx.fast.scaled_dot_product_attention(
query, key, value, scale=scale, mask=mask
)
)
def silu(input, inplace=False):
if inplace:
unsupported("torchmlx.nn.functional.silu with inplace=True")
return input * mx.sigmoid(input)
return wrap(input * mx.sigmoid(input))
def softmax(input, dim=None, dtype=None):
value = input.astype(dtype) if dtype is not None else input
return mx.softmax(value, axis=dim)
return wrap(mx.softmax(value, axis=dim))
def __getattr__(name):
unsupported(f"torchmlx.nn.functional.{name}")
+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
+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 },
]
[[package]]
name = "nodeenv"
version = "1.10.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/24/bf/d1bda4f6168e0b2e9e5958945e01910052158313224ada5ce1fb2e1113b8/nodeenv-1.10.0.tar.gz", hash = "sha256:996c191ad80897d076bdfba80a41994c2b47c68e224c542b48feba42ba00f8bb", size = 55611 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/88/b2/d0896bdcdc8d28a7fc5717c305f1a861c26e18c05047949fb371034d98bd/nodeenv-1.10.0-py2.py3-none-any.whl", hash = "sha256:5bb13e3eed2923615535339b3c620e76779af4cb4c6a90deccc9e36b274d3827", size = 23438 },
]
[[package]]
name = "numpy"
version = "2.2.6"
@@ -686,6 +695,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878 },
]
[[package]]
name = "pyright"
version = "1.1.414"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nodeenv" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/e1/1b/244c7b710031ada80f27e579ec20d28a2285dfc318fed0339866b1047f12/pyright-1.1.414.tar.gz", hash = "sha256:523c0a97c60da6333234955c277730c9cf4f5bd6d5399e7b7d2b0fc5d3599524", size = 4154638 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/d7/ba/18b6e682ead424ad24bcc134339ae5d1b931cd9ae260540592a058a91279/pyright-1.1.414-py3-none-any.whl", hash = "sha256:2a6b4b3298c9eec174c5ed83bd338de6eee82df2992f3e1930e6199d381be36f", size = 6225049 },
]
[[package]]
name = "pytorchmlx"
version = "0.0.1"
@@ -698,6 +720,11 @@ dependencies = [
{ name = "torch" },
]
[package.dev-dependencies]
dev = [
{ name = "pyright" },
]
[package.metadata]
requires-dist = [
{ name = "mlx", specifier = ">=0.32.2,<0.33" },
@@ -705,6 +732,9 @@ requires-dist = [
{ name = "torch", specifier = ">=2.4,<3" },
]
[package.metadata.requires-dev]
dev = [{ name = "pyright", specifier = ">=1.1.414" }]
[[package]]
name = "setuptools"
version = "84.0.0"