Skip to content

Commit 97f93cf

Browse files
authored
MAINT: refactor library to split up large files (#928)
1 parent 9ba8b3a commit 97f93cf

55 files changed

Lines changed: 6059 additions & 5789 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

meson.build

Lines changed: 33 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -10,25 +10,45 @@ py = import('python').find_installation()
1010
sources = {
1111
'array_api_extra': files(
1212
'src/array_api_extra/__init__.py',
13-
'src/array_api_extra/_delegation.py',
13+
'src/array_api_extra/_at.py',
14+
'src/array_api_extra/_creation.py',
15+
'src/array_api_extra/_elementwise.py',
16+
'src/array_api_extra/_indexing.py',
17+
'src/array_api_extra/_lazy.py',
18+
'src/array_api_extra/_linalg.py',
19+
'src/array_api_extra/_manipulation.py',
20+
'src/array_api_extra/_searching.py',
21+
'src/array_api_extra/_set.py',
22+
'src/array_api_extra/_sorting.py',
23+
'src/array_api_extra/_statistical.py',
1424
'src/array_api_extra/py.typed',
15-
'src/array_api_extra/testing.py',
25+
),
26+
'array_api_extra/_agnostic': files(
27+
'src/array_api_extra/_agnostic/__init__.py',
28+
'src/array_api_extra/_agnostic/_creation.py',
29+
'src/array_api_extra/_agnostic/_elementwise.py',
30+
'src/array_api_extra/_agnostic/_indexing.py',
31+
'src/array_api_extra/_agnostic/_inspection.py',
32+
'src/array_api_extra/_agnostic/_linalg.py',
33+
'src/array_api_extra/_agnostic/_manipulation.py',
34+
'src/array_api_extra/_agnostic/_searching.py',
35+
'src/array_api_extra/_agnostic/_set.py',
36+
'src/array_api_extra/_agnostic/_sorting.py',
37+
'src/array_api_extra/_agnostic/_statistical.py',
38+
),
39+
'array_api_extra/testing': files(
40+
'src/array_api_extra/testing/__init__.py',
41+
'src/array_api_extra/testing/_testing.py',
1642
),
1743
'array_api_extra/_lib': files(
1844
'src/array_api_extra/_lib/__init__.py',
19-
'src/array_api_extra/_lib/_at.py',
2045
'src/array_api_extra/_lib/_backends.py',
21-
'src/array_api_extra/_lib/_funcs.py',
22-
'src/array_api_extra/_lib/_lazy.py',
46+
'src/array_api_extra/_lib/_compat.py',
47+
'src/array_api_extra/_lib/_compat.pyi',
48+
'src/array_api_extra/_lib/_helpers.py',
2349
'src/array_api_extra/_lib/_testing.py',
24-
),
25-
'array_api_extra/_lib/_utils': files(
26-
'src/array_api_extra/_lib/_utils/__init__.py',
27-
'src/array_api_extra/_lib/_utils/_compat.py',
28-
'src/array_api_extra/_lib/_utils/_compat.pyi',
29-
'src/array_api_extra/_lib/_utils/_helpers.py',
30-
'src/array_api_extra/_lib/_utils/_typing.py',
31-
'src/array_api_extra/_lib/_utils/_typing.pyi',
50+
'src/array_api_extra/_lib/_typing.py',
51+
'src/array_api_extra/_lib/_typing.pyi',
3252
),
3353
}
3454

src/array_api_extra/__init__.py

Lines changed: 13 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -1,46 +1,22 @@
11
"""Extra array functions built on top of the array API standard."""
22

33
from . import testing
4-
from ._delegation import (
5-
argpartition,
6-
atleast_nd,
7-
broadcast_shapes,
8-
cov,
9-
create_diagonal,
10-
deg2rad,
11-
diag_indices,
12-
expand_dims,
13-
isclose,
14-
isin,
15-
kron,
16-
nan_to_num,
17-
nanmax,
18-
nanmin,
19-
nansum,
20-
nunique,
21-
one_hot,
22-
pad,
23-
partition,
24-
rad2deg,
25-
searchsorted,
26-
setdiff1d,
27-
sinc,
28-
tril_indices,
29-
triu_indices,
30-
union1d,
31-
unravel_index,
32-
)
33-
from ._lib._at import at
34-
from ._lib._funcs import (
35-
angle,
36-
apply_where,
37-
default_dtype,
38-
)
39-
from ._lib._lazy import lazy_apply
4+
from ._agnostic._elementwise import angle, apply_where
5+
from ._agnostic._inspection import default_dtype
6+
from ._at import at
7+
from ._creation import create_diagonal, one_hot
8+
from ._elementwise import deg2rad, isclose, nan_to_num, rad2deg, sinc
9+
from ._indexing import diag_indices, tril_indices, triu_indices, unravel_index
10+
from ._lazy import lazy_apply
11+
from ._linalg import kron
12+
from ._manipulation import atleast_nd, broadcast_shapes, expand_dims, pad
13+
from ._searching import searchsorted
14+
from ._set import isin, nunique, setdiff1d, union1d
15+
from ._sorting import argpartition, partition
16+
from ._statistical import cov, nanmax, nanmin, nansum
4017

4118
__version__ = "0.11.2.dev0"
4219

43-
# pylint: disable=duplicate-code
4420
__all__ = [
4521
"__version__",
4622
"angle",
Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,27 @@
1+
"""Array-agnostic function implementations."""
2+
3+
from . import (
4+
_creation,
5+
_elementwise,
6+
_indexing,
7+
_inspection,
8+
_linalg,
9+
_manipulation,
10+
_searching,
11+
_set,
12+
_sorting,
13+
_statistical,
14+
)
15+
16+
__all__ = [
17+
"_creation",
18+
"_elementwise",
19+
"_indexing",
20+
"_inspection",
21+
"_linalg",
22+
"_manipulation",
23+
"_searching",
24+
"_set",
25+
"_sorting",
26+
"_statistical",
27+
]
Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
"""Array-agnostic implementations for creation functions."""
2+
3+
from .._at import at
4+
from .._lib import _compat
5+
from .._lib._helpers import eager_shape, ndindex
6+
from .._lib._typing import Array, ArrayNamespace
7+
8+
__all__ = ["create_diagonal", "one_hot"]
9+
10+
11+
def create_diagonal(
12+
x: Array, /, *, offset: int = 0, xp: ArrayNamespace
13+
) -> Array: # numpydoc ignore=PR01,RT01
14+
"""See docstring in array_api_extra._delegation."""
15+
x_shape = eager_shape(x)
16+
batch_dims = x_shape[:-1]
17+
n = x_shape[-1] + abs(offset)
18+
diag = xp.zeros((*batch_dims, n**2), dtype=x.dtype, device=_compat.device(x))
19+
20+
target_slice = slice(
21+
offset if offset >= 0 else abs(offset) * n,
22+
min(n * (n - offset), diag.shape[-1]),
23+
n + 1,
24+
)
25+
for index in ndindex(*batch_dims):
26+
diag = at(diag)[(*index, target_slice)].set(x[(*index, slice(None))])
27+
return xp.reshape(diag, (*batch_dims, n, n))
28+
29+
30+
def one_hot(
31+
x: Array,
32+
/,
33+
num_classes: int,
34+
*,
35+
xp: ArrayNamespace,
36+
) -> Array: # numpydoc ignore=PR01,RT01
37+
"""See docstring in `array_api_extra._delegation.py`."""
38+
# TODO: Benchmark whether this is faster on the NumPy backend:
39+
# if is_numpy_array(x):
40+
# out = xp.zeros((x.size, num_classes), dtype=dtype)
41+
# out[xp.arange(x.size), xp.reshape(x, (-1,))] = 1
42+
# return xp.reshape(out, (*x.shape, num_classes))
43+
range_num_classes = xp.arange(num_classes, dtype=x.dtype, device=_compat.device(x))
44+
return x[..., xp.newaxis] == range_num_classes

0 commit comments

Comments
 (0)