Skip to content
Merged
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
12 changes: 1 addition & 11 deletions decent_array/_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ def __init__(self, value: ArrayTypes) -> None:
"Call set_backend() to initialize the interoperability layer."
)

self.value: Any = value
self.value: Any = _BACKEND_INSTANCE.native_asarray(value)
self._backend: Backend = _BACKEND_INSTANCE

# Binary arithmetic ----------------------------------------------------
Expand Down Expand Up @@ -445,16 +445,6 @@ def mT(self) -> Array: # noqa: N802
"""Return the matrix transpose (last two dimensions swapped)."""
return self._backend.matrix_transpose(self)

@property
def any(self) -> bool:
"""Return True if any element of the array is truthy."""
return self._backend.any(self)

@property
def all(self) -> bool:
"""Return True if all elements of the array are truthy."""
return self._backend.all(self)

@property
def device(self) -> Devices:
"""Return the device of the array."""
Expand Down
28 changes: 21 additions & 7 deletions decent_array/interoperability/_abstracts/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from numpy.typing import NDArray

from decent_array import Array
from decent_array.types import ArrayKey, ArrayTypes
from decent_array.types import ArrayKey, ArrayLike, ArrayTypes
from decent_array.types._dtypes import dtype


Expand Down Expand Up @@ -88,9 +88,23 @@ def from_numpy(self, x: NDArray[Any]) -> Array:
def from_numpy_like(self, x: NDArray[Any], like: Array) -> Array:
"""Convert a Numpy array to an :class:`Array` on this backend, matching shape and type of ``like``."""

def asarray(self, x: ArrayTypes | Array) -> Array:
"""
Convert `x` into an :class:`~decent_array.Array` on the active backend.

`x` can be a scalar, a backend-native array/tensor, or :class:`~decent_array.Array`, in which case it acts as
the identity.

"""
from decent_array._array import Array # noqa: PLC0415

if isinstance(x, Array):
return x
return Array(self.native_asarray(x))

@abstractmethod
def asarray(self, x: bool | int | float | complex) -> Array:
"""Convert a Python scalar to an :class:`Array` on this backend."""
def native_asarray(self, x: ArrayTypes) -> ArrayLike:
"""Wrap the backend-native asarray operation."""

@abstractmethod
def to_scalar(self, x: Array) -> Any: # noqa: ANN401
Expand Down Expand Up @@ -191,12 +205,12 @@ def max(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: boo
"""Maximum of ``x`` along ``axis``."""

@abstractmethod
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
"""Return True if any element of ``x`` is truthy."""
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
"""Test whether any array element along axis is truthy."""

@abstractmethod
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
"""Return True if all elements of ``x`` are truthy."""
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
"""Test whether all array elements along axis are truthy."""

# Math elementwise — both operands may be Array or scalar (operator dunders pass
# either). ``Array | float`` covers both because PEP 484's numeric tower implicitly
Expand Down
11 changes: 9 additions & 2 deletions decent_array/interoperability/_iop/manipulations.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,8 +66,15 @@ def from_numpy_like(x: NDArray[Any], like: Array) -> Array:
return _BACKEND_INSTANCE.from_numpy_like(x, like)


def asarray(x: float | bool) -> Array:
"""Convert a Python scalar to an :class:`~decent_array.Array` on the active backend."""
def asarray(x: ArrayTypes | Array) -> Array:
"""
Convert `x` into an :class:`~decent_array.Array` on the active backend.

`x` can be a scalar, a native array, a nested sequence, an object supporting Python's buffer protocol. Note that `x`
is passed directly to the equivalent backend-native asarray function. If `x` is an :class:`~decent_array.Array`
instance, it is returned unchanged.

"""
if _BACKEND_INSTANCE is None:
raise no_backend_error
return _BACKEND_INSTANCE.asarray(x)
Expand Down
8 changes: 4 additions & 4 deletions decent_array/interoperability/_iop/reductions.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,15 +77,15 @@ def max( # noqa: A001
return _BACKEND_INSTANCE.max(x, axis, keepdims)


def any(x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool: # noqa: A001
"""Return True if any element of ``x`` is truthy."""
def any(x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array: # noqa: A001
"""Test whether any array element along axis is truthy."""
if _BACKEND_INSTANCE is None:
raise no_backend_error
return _BACKEND_INSTANCE.any(x, axis, keepdims)


def all(x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool: # noqa: A001
"""Return True if all elements of ``x`` are truthy."""
def all(x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array: # noqa: A001
"""Test whether all array elements along axis are truthy."""
if _BACKEND_INSTANCE is None:
raise no_backend_error
return _BACKEND_INSTANCE.all(x, axis, keepdims)
15 changes: 8 additions & 7 deletions decent_array/interoperability/_jax/jax_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
from decent_array._utils import is_scalar, unwrap
from decent_array.interoperability._abstracts import Backend
from decent_array.interoperability._backend_manager import register_backend
from decent_array.types import ArrayKey, ArrayTypes, Devices, Frameworks
from decent_array.types import ArrayKey, ArrayLike, ArrayTypes, Devices, Frameworks
from decent_array.types._dtypes import dtype


Expand Down Expand Up @@ -91,8 +91,9 @@ def from_numpy_like(self, x: NDArray[Any], like: Array) -> Array:
v = like.value
return Array(jnp.asarray(x, dtype=v.dtype, device=v.device))

def asarray(self, x: bool | int | float | complex) -> Array:
return Array(jnp.array(x, device=self._native_device))
def native_asarray(self, x: ArrayTypes) -> ArrayLike:
"""Wrap the backend-native asarray operation."""
return jnp.asarray(x, device=self._native_device)

def to_scalar(self, x: Array) -> Any: # noqa: ANN401
"""
Expand Down Expand Up @@ -187,11 +188,11 @@ def min(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: boo
def max(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(jnp.max(x.value, axis=axis, keepdims=keepdims))

def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(jnp.any(x.value, axis=axis, keepdims=keepdims))
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(jnp.any(x.value, axis=axis, keepdims=keepdims))

def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(jnp.all(x.value, axis=axis, keepdims=keepdims))
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(jnp.all(x.value, axis=axis, keepdims=keepdims))

# Math elementwise — JAX arrays are immutable; "in-place" ops rebind the wrapper.
# Operands may be Array or scalar (operator dunders pass either); ``Array | float``
Expand Down
19 changes: 10 additions & 9 deletions decent_array/interoperability/_numpy/numpy_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
from decent_array._utils import is_scalar, unwrap
from decent_array.interoperability._abstracts import Backend
from decent_array.interoperability._backend_manager import register_backend
from decent_array.types import ArrayKey, ArrayTypes, Devices, Frameworks
from decent_array.types import ArrayKey, ArrayLike, ArrayTypes, Devices, Frameworks
from decent_array.types._dtypes import dtype


Expand Down Expand Up @@ -84,8 +84,9 @@ def from_numpy_like(self, x: NDArray[Any], like: Array) -> Array:
# NumPy has no device dimension, so only the dtype of ``like`` matters.
return Array(np.asarray(x, dtype=like.value.dtype))

def asarray(self, x: bool | int | float | complex) -> Array:
return Array(np.array(x))
def native_asarray(self, x: ArrayTypes) -> ArrayLike:
"""Wrap the backend-native asarray operation."""
return np.asarray(x)

def to_scalar(self, x: Array) -> Any: # noqa: ANN401
"""
Expand Down Expand Up @@ -194,17 +195,17 @@ def max(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: boo
return Array(np.max(v, axis=axis, keepdims=True))
return Array(np.max(v, axis=axis))

def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
v = cast("np.ndarray[Any, Any]", x.value)
if keepdims:
return bool(np.any(v, axis=axis, keepdims=True))
return bool(np.any(v, axis=axis))
return Array(np.any(v, axis=axis, keepdims=True))
return Array(np.any(v, axis=axis, keepdims=False))

def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
v = cast("np.ndarray[Any, Any]", x.value)
if keepdims:
return bool(np.all(v, axis=axis, keepdims=True))
return bool(np.all(v, axis=axis))
return Array(np.all(v, axis=axis, keepdims=True))
return Array(np.all(v, axis=axis, keepdims=False))

# Math elementwise — operands may be Array or scalar (operator dunders pass either).
# ``Array | float`` covers both: PEP 484's numeric tower implicitly admits ``int``.
Expand Down
15 changes: 8 additions & 7 deletions decent_array/interoperability/_pytorch/pytorch_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from decent_array._utils import is_scalar, unwrap
from decent_array.interoperability._abstracts import Backend
from decent_array.interoperability._backend_manager import register_backend
from decent_array.types import ArrayKey, ArrayTypes, Devices, Frameworks
from decent_array.types import ArrayKey, ArrayLike, ArrayTypes, Devices, Frameworks
from decent_array.types._dtypes import dtype


Expand Down Expand Up @@ -88,8 +88,9 @@ def from_numpy_like(self, x: NDArray[Any], like: Array) -> Array:
v = like.value
return Array(torch.from_numpy(x).to(dtype=v.dtype, device=v.device))

def asarray(self, x: bool | int | float | complex) -> Array:
return Array(torch.tensor(x, device=self._native_device))
def native_asarray(self, x: ArrayTypes) -> ArrayLike:
"""Wrap the backend-native asarray operation."""
return torch.as_tensor(x, device=self._native_device)

def to_scalar(self, x: Array) -> Any: # noqa: ANN401
"""
Expand Down Expand Up @@ -201,11 +202,11 @@ def max(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: boo
return Array(torch.max(v))
return Array(torch.amax(v, dim=axis, keepdim=keepdims))

def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(torch.any(x.value, dim=axis, keepdim=keepdims).item())
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(torch.any(x.value, dim=axis, keepdim=keepdims))

def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(torch.all(x.value, dim=axis, keepdim=keepdims).item())
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(torch.all(x.value, dim=axis, keepdim=keepdims))

# Math elementwise — operands may be Array or scalar (operator dunders pass either).
# ``Array | float`` covers both: PEP 484's numeric tower implicitly admits ``int``.
Expand Down
19 changes: 10 additions & 9 deletions decent_array/interoperability/_tensorflow/tensorflow_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
from decent_array._utils import is_scalar, unwrap
from decent_array.interoperability._abstracts import Backend
from decent_array.interoperability._backend_manager import register_backend
from decent_array.types import ArrayKey, ArrayTypes, Devices, Frameworks
from decent_array.types import ArrayKey, ArrayLike, ArrayTypes, Devices, Frameworks
from decent_array.types._dtypes import dtype


Expand Down Expand Up @@ -98,11 +98,12 @@ def from_numpy_like(self, x: NDArray[Any], like: Array) -> Array:
with tf.device(v.device):
return Array(tf.convert_to_tensor(x, dtype=v.dtype))

def asarray(self, x: bool | int | float | complex) -> Array:
"""Convert a Python scalar to an :class:`Array` on this backend."""
def native_asarray(self, x: ArrayTypes) -> ArrayLike:
"""Wrap the backend-native asarray operation."""
# TensorFlow's type stubs do not cover all inputs accepted by
# convert_to_tensor at runtime.
with tf.device(self._native_device):
# Its not a tf tensor but mypyc doesn't import tf so it complains about unsude type-ignores
return Array(tf.convert_to_tensor(cast("tf.Tensor", x)))
return tf.convert_to_tensor(cast("Any", x))

def to_scalar(self, x: Array) -> Any: # noqa: ANN401
"""
Expand Down Expand Up @@ -210,11 +211,11 @@ def min(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: boo
def max(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(tf.reduce_max(x.value, axis=axis, keepdims=keepdims))

def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(tf.reduce_any(tf.cast(x.value, tf.bool), axis=axis, keepdims=keepdims).numpy())
def any(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(tf.reduce_any(tf.cast(x.value, tf.bool), axis=axis, keepdims=keepdims))

def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> bool:
return bool(tf.reduce_all(tf.cast(x.value, tf.bool), axis=axis, keepdims=keepdims).numpy())
def all(self, x: Array, axis: int | tuple[int, ...] | None = None, keepdims: bool = False) -> Array:
return Array(tf.reduce_all(tf.cast(x.value, tf.bool), axis=axis, keepdims=keepdims))

# Math elementwise — TF Tensors are immutable; "in-place" ops rebind the wrapper.
# Operands may be Array or scalar (operator dunders pass either); ``Array | float``
Expand Down
12 changes: 11 additions & 1 deletion decent_array/types/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from __future__ import annotations

from collections.abc import Buffer, Sequence
from enum import Enum
from typing import TYPE_CHECKING, SupportsIndex, TypeAlias, Union

Expand All @@ -14,14 +15,23 @@

from decent_array._array import Array

Scalar: TypeAlias = bool | int | float | complex | numpy.generic # noqa: UP040
"""
Type alias for scalar types supported in decent-array.
"""

ArrayLike: TypeAlias = Union["numpy.ndarray", "torch.Tensor", "tf.Tensor", "jax.Array"] # noqa: UP040
"""
Type alias for array-like types supported in decent-array, including NumPy arrays,
PyTorch tensors, TensorFlow tensors, and JAX arrays.
"""

ArrayTypes: TypeAlias = bool | int | float | complex | numpy.generic | ArrayLike # noqa: UP040
NestedSequence: TypeAlias = "Sequence[Scalar | NestedSequence]" # noqa: UP040
"""
Type alias for nested sequences supported in decent-array.
"""

ArrayTypes: TypeAlias = Scalar | ArrayLike | NestedSequence | Buffer # noqa: UP040
"""
Type alias for supported scalar/array types in decent-array.
"""
Expand Down
1 change: 1 addition & 0 deletions docs/source/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
("py:class", "float64"),
("py:class", "numpy._typing._array_like._SupportsArray"),
("py:class", "numpy._typing._nested_sequence._NestedSequence"),
("py:class", "Sequence[Scalar | NestedSequence]"),
("py:class", "T"),
]

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "decent-array"
version = "0.2.5"
version = "0.2.6"
authors = [{name = "Simon Granström"}, {name = "Nicola Bastianello"}]
maintainers = [{name = "Team Decent"}]
description = "A library of array operations and linear algebra primitives for interoperability across ML frameworks."
Expand Down
47 changes: 27 additions & 20 deletions tests/test_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,33 @@ def test_init_records_active_backend(backend: tuple) -> None:
assert isinstance(arr, Array)


@pytest.mark.parametrize(
"value",
[
True,
1,
1.0,
1.0 + 2.0j,
np.int32(1),
np.float32(1.0),
np.complex64(1.0 + 2.0j),
],
)
def test_init_converts_to_backend_array(value, backend: tuple) -> None:
arr = Array(value)

assert isinstance(arr, Array)
assert arr.ndim == 0
assert arr.size == 1
assert arr.shape == ()


def test_init_scalar_value(backend: tuple) -> None:
arr = Array(4)

np.testing.assert_array_equal(iop.to_numpy(arr), np.array(4))


# Binary arithmetic -------------------------------------------------------


Expand Down Expand Up @@ -528,26 +555,6 @@ def test_mT_raises_for_rank_lt_2(backend: tuple) -> None:
_ = a.mT


def test_any_true(backend: tuple) -> None:
a = _create_array([0.0, 0.0, 1.0])
assert a.any is True


def test_any_false(backend: tuple) -> None:
a = _create_array([0.0, 0.0, 0.0])
assert a.any is False


def test_all_true(backend: tuple) -> None:
a = _create_array([1.0, 2.0, 3.0])
assert a.all is True


def test_all_false(backend: tuple) -> None:
a = _create_array([1.0, 0.0, 3.0])
assert a.all is False


def test_device_property(backend: tuple) -> None:
_framework, device = backend
a = iop.zeros((3,))
Expand Down
Loading
Loading