Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,8 @@ jobs:
no-coverage: true,
x64_runner: ubuntu-latest,
}
- environment: tests-mlx
runs-on: macos-15
- {environment: tests-run-deps, task: tests-run-deps-cov}
exclude:
- {environment: tests-numpy1, platform: windows} # data-apis/array-api-extra#901
Expand Down
13,278 changes: 7,237 additions & 6,041 deletions pixi.lock

Large diffs are not rendered by default.

8 changes: 8 additions & 0 deletions pixi.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ platforms = [
"linux-aarch64",
"osx-64",
"osx-arm64",
{ name = "osx-arm64-macos14", platform = "osx-arm64", macos = "14.5" },
"win-64",
{ platform = "linux-64", cuda = "12.9" },
{ platform = "win-64", cuda = "12.9" },
Expand Down Expand Up @@ -69,6 +70,7 @@ tests-backends = {
solve-group = "backends",
}
tests-backends-py311.features = ["py311", "tests", "backends"]
tests-mlx = { features = ["py314", "tests", "mlx"], solve-group = "mlx" }

# CUDA not available on free github actions and on some developers' PCs
dev-cuda = {
Expand Down Expand Up @@ -378,6 +380,12 @@ mparray = ">=0.2.2"
[feature.backends.target.unix.dependencies]
jax = ">=0.10.2" # waiting for conda-forge/jaxlib-feedstock#326

[feature.mlx.target.osx-arm64.pypi-dependencies]

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

0.32.3 is not yet available on conda-forge. See

mlx = ">=0.32.2"

[feature.mlx]
platforms = ["osx-arm64-macos14"]

# Backends that require a GPU host and a CUDA driver.
[feature.cuda-backends]
platforms = ["linux-64-cuda-12-9", "win-64-cuda-12-9"]
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@ disable_error_code = ["no-any-return"]

[[tool.mypy.overrides]]
# slow or unavailable on Windows; do not add to the lint env
module = ["cupy.*", "jax.*", "sparse.*", "torch.*"]
module = ["cupy.*", "jax.*", "mlx.*", "sparse.*", "torch.*"]
ignore_missing_imports = true

[[tool.mypy.overrides]]
Expand Down
1 change: 1 addition & 0 deletions src/array_api_extra/_lib/_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ class Backend(Enum): # numpydoc ignore=PR02
NUMPY_READONLY = "numpy:readonly"
MPARRAY = "mparray"
CUPY = "cupy"
MLX = "mlx.core"
TORCH = "torch"
TORCH_GPU = "torch:gpu"
DASK = "dask.array"
Expand Down
159 changes: 152 additions & 7 deletions src/array_api_extra/_lib/_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,27 @@
# Allow packages that vendor both `array-api-extra` and
# `array-api-compat` to override the import location

# pylint: disable=duplicate-code
from functools import cache
from types import ModuleType
from typing import TYPE_CHECKING, ClassVar

if TYPE_CHECKING: # pragma: no cover
from typing_extensions import override
else:

def override(func):
return func


# pylint: disable=duplicate-code,redefined-outer-name
try:
from ..._array_api_compat_vendor import (
array_namespace,
device,
array_namespace as _array_namespace,
)
from ...._array_api_compat_vendor import (
device as _device,
)
from ...._array_api_compat_vendor import (
is_array_api_obj,
is_array_api_strict_namespace,
is_cupy_array,
Expand All @@ -24,12 +40,18 @@
is_torch_namespace,
is_writeable_array,
size,
to_device,
)
from ...._array_api_compat_vendor import (
to_device as _to_device,
)
except ImportError:
from array_api_compat import (
array_namespace,
device,
array_namespace as _array_namespace,
)
from array_api_compat import (
device as _device,
)
from array_api_compat import (
is_array_api_obj,
is_array_api_strict_namespace,
is_cupy_array,
Expand All @@ -47,8 +69,131 @@
is_torch_namespace,
is_writeable_array,
size,
to_device,
)
from array_api_compat import (
to_device as _to_device,
)


class _MLXNamespaceInfo:
def __init__(self, info: object) -> None:
self._info = info

Check failure on line 80 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Type annotation for attribute `_info` is required because this class is not decorated with `@final` (reportUnannotatedClassAttribute)

def __getattr__(self, name: str) -> object:
return getattr(self._info, name)

def default_dtypes(self, *, device: object = None) -> object:
_ = device
return self._info.default_dtypes() # type: ignore[attr-defined]

Check failure on line 87 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Cannot access attribute "default_dtypes" for class "object"   Attribute "default_dtypes" is unknown (reportAttributeAccessIssue)
Comment on lines +78 to +87

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Most of the changes in these files should be in array-api-compat



class _MLXNamespace(ModuleType):
_device_functions: ClassVar[set[str]] = {
"arange",
"asarray",
"empty",
"eye",
"full",
"linspace",
"ones",
"zeros",
}

def __init__(self, namespace: ModuleType) -> None:
super().__init__(namespace.__name__)
self._namespace = namespace

Check failure on line 104 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Type annotation for attribute `_namespace` is required because this class is not decorated with `@final` (reportUnannotatedClassAttribute)

@override
def __getattr__(self, name: str) -> object:
if name == "__array_api_version__":
return "2024.12"
if name == "bool":
return self._namespace.bool_
if name == "__array_namespace_info__":
info = getattr(self._namespace, name)
return lambda: _MLXNamespaceInfo(info())
if name == "signbit":
return lambda x: x < 0
function = getattr(self._namespace, name)
if name not in {"argsort", "astype", "result_type", "sort"} and (
name not in self._device_functions
):
return function

def compatible_call(
*args: object,
device: object = None,
copy: bool | None = None,
**kwargs: object,
) -> object:
if name == "full" and "fill_value" in kwargs:
args = (*args, kwargs.pop("fill_value"))
if name == "asarray" and copy is not None:
kwargs["copy"] = copy
if name == "asarray" and "dtype" not in kwargs:
input_array = args[0] if args else kwargs.get("a")
input_dtype = getattr(input_array, "dtype", None)
dtype_name = getattr(input_dtype, "name", None)
if dtype_name is not None and hasattr(self._namespace, dtype_name):
kwargs["dtype"] = getattr(self._namespace, dtype_name)
if name == "astype":
_ = copy
return function(*args, **kwargs)
if name in {"argsort", "sort"}:
kwargs.pop("stable", None)

Check failure on line 143 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Result of call expression is of type "object" and is not used; assign to variable "_" if this is intentional (reportUnusedCallResult)
if name == "result_type":
import mlx.core as mx

args = tuple(
mx.asarray(value)

Check failure on line 148 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Argument type is partially unknown   Argument corresponds to parameter "iterable" in function "__new__"   Argument type is "Generator[Unknown | object, None, None]" (reportUnknownArgumentType)
if isinstance(value, int | float | complex | bool)
else value
for value in args
)
if device is None:
return function(*args, **kwargs)
import mlx.core as mx

with mx.stream(device):
return function(*args, **kwargs)

return compatible_call


@cache
def _wrap_mlx_namespace(namespace: ModuleType) -> ModuleType:
return _MLXNamespace(namespace)


def array_namespace(
*xs: object, api_version: str | None = None, use_compat: bool | None = None
) -> ModuleType:
namespace = _array_namespace(*xs, api_version=api_version, use_compat=use_compat)
if namespace.__name__ == "mlx.core":
return _wrap_mlx_namespace(namespace)

Check failure on line 173 in src/array_api_extra/_lib/_compat.py

View workflow job for this annotation

GitHub Actions / Lint

Argument type is partially unknown   Argument corresponds to parameter "args" in function "__call__"   Argument type is "ModuleType | Unknown" (reportUnknownArgumentType)
return namespace


def device(x: object, /) -> object:
if type(x).__module__.startswith("mlx."):
import mlx.core as mx

return mx.default_device()
return _device(x)


def to_device(x: object, device: object, /, *, stream: object = None) -> object:
if type(x).__module__.startswith("mlx."):
import mlx.core as mx

if device == "cpu":
device = mx.cpu
elif device == "gpu":
device = mx.gpu
with mx.stream(device):
return mx.array(x)
return _to_device(x, device, stream=stream)


__all__ = [
"array_namespace",
Expand Down
16 changes: 15 additions & 1 deletion tests/main/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,14 @@ def xp(
_setup_jax(library)
elif library.like(Backend.TORCH):
_setup_torch(library)
elif library is Backend.MLX:
_setup_mlx()

import mlx.core as mx

with mx.stream(mx.cpu), patch_lazy_xp_functions(request, xp=xp):
yield xp
return

# On Dask and JAX, monkey-patch all functions tagged by `lazy_xp_function`
# in the global scope of the module containing the test function.
Expand Down Expand Up @@ -189,6 +197,12 @@ def _setup_torch(library: Backend) -> None:
torch.set_default_device("cpu")


def _setup_mlx() -> None:
import mlx.core as mx

mx.set_default_device(mx.cpu)


# Can select the test with `pytest -k dask`
@pytest.fixture(params=[Backend.DASK.pytest_param()])
def da(
Expand Down Expand Up @@ -244,6 +258,6 @@ def device(
@pytest.fixture
def infinity(library: Backend) -> float:
"""Retrieve the positive infinity value for the given backend."""
if library in (Backend.TORCH, Backend.TORCH_GPU):
if library in (Backend.MLX, Backend.TORCH, Backend.TORCH_GPU):
return 3.4028235e38
return 1.7976931348623157e308
17 changes: 16 additions & 1 deletion tests/main/test_at.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,7 @@ def assert_copy(
)
def test_update_ops(
xp: ArrayNamespace,
library: Backend,
copy: bool | None,
op: _AtOp,
y: float,
Expand All @@ -129,6 +130,8 @@ def test_update_ops(
x_ndim: int,
y_ndim: int,
):
if bool_mask and library is Backend.MLX:
pytest.skip("MLX does not support boolean indexing")
if x_ndim == 1:
x = xp.asarray([10.0, 20.0, 30.0])
idx = xp.asarray([False, True, True]) if bool_mask else slice(1, None)
Expand All @@ -153,6 +156,7 @@ def test_update_ops(


@pytest.mark.parametrize("op", list(_AtOp))
@pytest.mark.skip_xp_backend(Backend.MLX, reason="boolean indexing is unsupported")
def test_copy_default(xp: ArrayNamespace, library: Backend, op: _AtOp):
"""
Test that the default copy behaviour is False for writeable arrays
Expand Down Expand Up @@ -235,6 +239,8 @@ def test_incompatible_dtype(
UFuncTypeError: Cannot cast ufunc 'divide' output from dtype('float64')
to dtype('int64') with casting rule 'same_kind'
"""
if bool_mask and library is Backend.MLX:
pytest.skip("MLX does not support boolean indexing")
x = xp.asarray([2, 4])
idx = xp.asarray([True, False]) if bool_mask else slice(None)
z = None
Expand All @@ -254,6 +260,13 @@ def test_incompatible_dtype(
with pytest.raises(Exception, match=r"cast|promote|dtype"):
_ = at_op(x, idx, op, 1.1, copy=copy)

elif library is Backend.MLX:
if op is _AtOp.DIVIDE:
with pytest.raises(Exception, match="cast"):
_ = at_op(x, idx, op, 1.1, copy=copy)
else:
z = at_op(x, idx, op, 1.1, copy=copy)

elif op in (_AtOp.SET, _AtOp.MIN, _AtOp.MAX):
# There is no __i<op>__ version of min/max.
# libraries other than array-api-strict are happy with
Expand All @@ -277,7 +290,9 @@ def test_bool_mask_nd(xp: ArrayNamespace):


@pytest.mark.parametrize("bool_mask", [False, True])
def test_no_inf_warnings(xp: ArrayNamespace, bool_mask: bool):
def test_no_inf_warnings(xp: ArrayNamespace, library: Backend, bool_mask: bool):
if bool_mask and library is Backend.MLX:
pytest.skip("MLX does not support boolean indexing")
x = xp.asarray([math.inf, 1.0, 2.0])
idx = ~xp.isinf(x) if bool_mask else slice(1, None)
# inf - inf -> nan with a warning
Expand Down
3 changes: 3 additions & 0 deletions tests/main/test_creation.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,9 @@ def test_2d(self, xp: ArrayNamespace):
@pytest.mark.skip_xp_backend(
Backend.ARRAY_API_STRICTEST, reason="backend doesn't support Boolean indexing"
)
@pytest.mark.skip_xp_backend(
Backend.MLX, reason="backend doesn't support Boolean indexing"
)
def test_abstract_size(self, xp: ArrayNamespace):
x = xp.arange(5)
x = x[x > 2]
Expand Down
10 changes: 10 additions & 0 deletions tests/main/test_elementwise.py
Original file line number Diff line number Diff line change
Expand Up @@ -398,6 +398,9 @@ def test_bool_dtype(self, xp: ArrayNamespace):

@pytest.mark.skip_xp_backend(Backend.SPARSE, reason="index by sparse array")
@pytest.mark.skip_xp_backend(Backend.ARRAY_API_STRICTEST, reason="unknown shape")
@pytest.mark.skip_xp_backend(
Backend.MLX, reason="backend doesn't support Boolean indexing"
)
def test_none_shape(self, xp: ArrayNamespace):
a = xp.asarray([1, 5, 0])
b = xp.asarray([1, 4, 2])
Expand All @@ -407,6 +410,9 @@ def test_none_shape(self, xp: ArrayNamespace):

@pytest.mark.skip_xp_backend(Backend.SPARSE, reason="index by sparse array")
@pytest.mark.skip_xp_backend(Backend.ARRAY_API_STRICTEST, reason="unknown shape")
@pytest.mark.skip_xp_backend(
Backend.MLX, reason="backend doesn't support Boolean indexing"
)
def test_none_shape_bool(self, xp: ArrayNamespace):
a = xp.asarray([True, True, False])
b = xp.asarray([True, False, True])
Expand Down Expand Up @@ -604,6 +610,9 @@ def test_dtype(self, xp: ArrayNamespace, x: complex):
with pytest.raises(ValueError, match="real floating data type"):
_ = sinc(xp.asarray(x))

@pytest.mark.skip_xp_backend(
Backend.MLX, reason="float32 precision is insufficient for this assertion"
)
def test_3d(self, xp: ArrayNamespace):
x = np.arange(18, dtype=np.float64).reshape((3, 3, 2))
expected = np.zeros_like(x)
Expand All @@ -627,6 +636,7 @@ def test_simple(self, xp: ArrayNamespace):
expected = xp.asarray([0.0, 0.0], dtype=res.dtype)
assert_equal(res, expected)

@pytest.mark.skip_xp_backend(Backend.MLX, reason="backend has no complex128 dtype")
def test_basic(self, xp: ArrayNamespace):
x = xp.asarray(
[
Expand Down
Loading
Loading