haliax 1.4.dev382__tar.gz → 1.4.dev388__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- haliax-1.4.dev388/.playbooks/wrap-non-named.md +49 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/AGENTS.md +1 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/PKG-INFO +1 -1
- haliax-1.4.dev388/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/__init__.py +20 -1
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/ops.py +249 -1
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_ops.py +157 -0
- haliax-1.4.dev382/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.coveragerc +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.flake8 +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.gitignore +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/LICENSE +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/README.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/api.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/css/material.css +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/faq.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/fp8.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/index.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/indexing.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/matmul.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/nn.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/partitioning.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/rearrange.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/requirements.txt +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/scan.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/state-dict.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/tutorial.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/typing.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/docs/vmap.md +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/mkdocs.yml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/pyproject.toml +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/core.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/random.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/types.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/util.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/core_test.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_attention.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_axis.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_conv.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_debug.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_dot.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_hof.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_int8.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_nn.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_pool.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_random.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_scan.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/tests/test_utils.py +0 -0
- {haliax-1.4.dev382 → haliax-1.4.dev388}/uv.lock +0 -0
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
# Wrapping Functions with NamedArray Support
|
|
2
|
+
|
|
3
|
+
This playbook explains how to convert a regular JAX function that works on unnamed arrays into a Haliax function that accepts `NamedArray` inputs and returns `NamedArray` outputs.
|
|
4
|
+
|
|
5
|
+
## When is wrapping needed?
|
|
6
|
+
Many JAX primitives only operate on regular arrays. To integrate them in Haliax you should provide a thin wrapper that handles axis metadata. Simple elementwise operations and reductions have helper utilities.
|
|
7
|
+
|
|
8
|
+
## Elemwise Unary
|
|
9
|
+
For a unary function that acts elementwise (e.g. `jnp.abs`):
|
|
10
|
+
|
|
11
|
+
```python
|
|
12
|
+
from haliax import wrap_elemwise_unary
|
|
13
|
+
|
|
14
|
+
def abs(a):
|
|
15
|
+
return wrap_elemwise_unary(jnp.abs, a)
|
|
16
|
+
```
|
|
17
|
+
|
|
18
|
+
This preserves axis order and dtype.
|
|
19
|
+
|
|
20
|
+
## Elemwise Binary
|
|
21
|
+
For binary operations (e.g. `jnp.add`), decorate a function with `wrap_elemwise_binary`:
|
|
22
|
+
|
|
23
|
+
```python
|
|
24
|
+
from haliax import wrap_elemwise_binary
|
|
25
|
+
|
|
26
|
+
@wrap_elemwise_binary
|
|
27
|
+
def add(x1, x2):
|
|
28
|
+
return jnp.add(x1, x2)
|
|
29
|
+
```
|
|
30
|
+
|
|
31
|
+
Broadcasting between `NamedArray`s is handled automatically.
|
|
32
|
+
|
|
33
|
+
## Reductions
|
|
34
|
+
Reductions require choosing axes to eliminate. Use `wrap_reduction_call`:
|
|
35
|
+
|
|
36
|
+
```python
|
|
37
|
+
from haliax import wrap_reduction_call
|
|
38
|
+
|
|
39
|
+
def sum(a, axis=None):
|
|
40
|
+
return wrap_reduction_call(jnp.sum, a, axis)
|
|
41
|
+
```
|
|
42
|
+
|
|
43
|
+
`axis` can be an `AxisSelector` or tuple. The wrapper returns a `NamedArray` with those axes removed.
|
|
44
|
+
|
|
45
|
+
## Harder Cases
|
|
46
|
+
Some functions need bespoke handling. For example `jnp.unique` returns several arrays and may change shape unpredictably. There is no generic helper, so you will need to manually map between `NamedArray` axes and the outputs. Use the lower level utilities in `haliax.wrap` for broadcasting and axis lookup.
|
|
47
|
+
|
|
48
|
+
## Testing
|
|
49
|
+
Add tests to ensure that named and unnamed calls produce the same results and that axis names are preserved or removed correctly.
|
|
@@ -17,6 +17,7 @@ repository. Follow these notes when implementing new features or fixing bugs.
|
|
|
17
17
|
## Playbook
|
|
18
18
|
|
|
19
19
|
- Adding Haliax-style tensor typing annotations are described in @.playbooks/add-types.md
|
|
20
|
+
- Wrapping standard JAX functions so they operate on `NamedArray` is explained in @.playbooks/wrap-non-named.md
|
|
20
21
|
|
|
21
22
|
## Code Style
|
|
22
23
|
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev388
|
|
4
4
|
Summary: Named Tensors for Legible Deep Learning in JAX
|
|
5
5
|
Project-URL: Homepage, https://github.com/stanford-crfm/haliax
|
|
6
6
|
Project-URL: Bug Tracker, https://github.com/stanford-crfm/haliax/issues/
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev388"
|
|
@@ -66,7 +66,21 @@ from .core import (
|
|
|
66
66
|
from .haxtyping import Named
|
|
67
67
|
from .hof import fold, map, scan, vmap
|
|
68
68
|
from .jax_utils import tree_checkpoint_name
|
|
69
|
-
from .ops import
|
|
69
|
+
from .ops import (
|
|
70
|
+
clip,
|
|
71
|
+
isclose,
|
|
72
|
+
pad_left,
|
|
73
|
+
pad,
|
|
74
|
+
trace,
|
|
75
|
+
tril,
|
|
76
|
+
triu,
|
|
77
|
+
unique,
|
|
78
|
+
unique_values,
|
|
79
|
+
unique_counts,
|
|
80
|
+
unique_inverse,
|
|
81
|
+
unique_all,
|
|
82
|
+
where,
|
|
83
|
+
)
|
|
70
84
|
from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
|
|
71
85
|
from .specialized_fns import top_k
|
|
72
86
|
from .types import Scalar
|
|
@@ -1034,6 +1048,11 @@ __all__ = [
|
|
|
1034
1048
|
"vmap",
|
|
1035
1049
|
"trace",
|
|
1036
1050
|
"where",
|
|
1051
|
+
"unique",
|
|
1052
|
+
"unique_values",
|
|
1053
|
+
"unique_counts",
|
|
1054
|
+
"unique_inverse",
|
|
1055
|
+
"unique_all",
|
|
1037
1056
|
"clip",
|
|
1038
1057
|
"tril",
|
|
1039
1058
|
"triu",
|
|
@@ -3,6 +3,9 @@ from typing import Mapping, Optional, Union
|
|
|
3
3
|
|
|
4
4
|
import jax
|
|
5
5
|
import jax.numpy as jnp
|
|
6
|
+
from jaxtyping import ArrayLike
|
|
7
|
+
|
|
8
|
+
import haliax
|
|
6
9
|
|
|
7
10
|
from .axis import Axis, AxisSelector, axis_name
|
|
8
11
|
from .core import NamedArray, NamedOrNumeric, broadcast_arrays, broadcast_arrays_and_return_axes, named
|
|
@@ -196,4 +199,249 @@ def raw_array_or_scalar(x: NamedOrNumeric):
|
|
|
196
199
|
return x
|
|
197
200
|
|
|
198
201
|
|
|
199
|
-
|
|
202
|
+
@typing.overload
|
|
203
|
+
def unique(
|
|
204
|
+
array: NamedArray, Unique: Axis, *, axis: AxisSelector | None = None, fill_value: ArrayLike | None = None
|
|
205
|
+
) -> NamedArray:
|
|
206
|
+
...
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
@typing.overload
|
|
210
|
+
def unique(
|
|
211
|
+
array: NamedArray,
|
|
212
|
+
Unique: Axis,
|
|
213
|
+
*,
|
|
214
|
+
return_index: typing.Literal[True],
|
|
215
|
+
axis: AxisSelector | None = None,
|
|
216
|
+
fill_value: ArrayLike | None = None,
|
|
217
|
+
) -> tuple[NamedArray, NamedArray]:
|
|
218
|
+
...
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
@typing.overload
|
|
222
|
+
def unique(
|
|
223
|
+
array: NamedArray,
|
|
224
|
+
Unique: Axis,
|
|
225
|
+
*,
|
|
226
|
+
return_inverse: typing.Literal[True],
|
|
227
|
+
axis: AxisSelector | None = None,
|
|
228
|
+
fill_value: ArrayLike | None = None,
|
|
229
|
+
) -> tuple[NamedArray, NamedArray]:
|
|
230
|
+
...
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
@typing.overload
|
|
234
|
+
def unique(
|
|
235
|
+
array: NamedArray,
|
|
236
|
+
Unique: Axis,
|
|
237
|
+
*,
|
|
238
|
+
return_counts: typing.Literal[True],
|
|
239
|
+
axis: AxisSelector | None = None,
|
|
240
|
+
fill_value: ArrayLike | None = None,
|
|
241
|
+
) -> tuple[NamedArray, NamedArray]:
|
|
242
|
+
...
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
@typing.overload
|
|
246
|
+
def unique(
|
|
247
|
+
array: NamedArray,
|
|
248
|
+
Unique: Axis,
|
|
249
|
+
*,
|
|
250
|
+
return_index: bool = False,
|
|
251
|
+
return_inverse: bool = False,
|
|
252
|
+
return_counts: bool = False,
|
|
253
|
+
axis: AxisSelector | None = None,
|
|
254
|
+
fill_value: ArrayLike | None = None,
|
|
255
|
+
) -> NamedArray | tuple[NamedArray, ...]:
|
|
256
|
+
...
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def unique(
|
|
260
|
+
array: NamedArray,
|
|
261
|
+
Unique: Axis,
|
|
262
|
+
*,
|
|
263
|
+
return_index: bool = False,
|
|
264
|
+
return_inverse: bool = False,
|
|
265
|
+
return_counts: bool = False,
|
|
266
|
+
axis: AxisSelector | None = None,
|
|
267
|
+
fill_value: ArrayLike | None = None,
|
|
268
|
+
) -> NamedArray | tuple[NamedArray, ...]:
|
|
269
|
+
"""
|
|
270
|
+
Like jnp.unique, but with named axes.
|
|
271
|
+
|
|
272
|
+
Args:
|
|
273
|
+
array: The input array.
|
|
274
|
+
Unique: The name of the axis that will be created to hold the unique values.
|
|
275
|
+
fill_value: The value to use for the fill_value argument of jnp.unique
|
|
276
|
+
axis: The axis along which to find unique values.
|
|
277
|
+
return_index: If True, return the indices of the unique values.
|
|
278
|
+
return_inverse: If True, return the indices of the input array that would reconstruct the unique values.
|
|
279
|
+
"""
|
|
280
|
+
size = Unique.size
|
|
281
|
+
|
|
282
|
+
is_multireturn = return_index or return_inverse or return_counts
|
|
283
|
+
|
|
284
|
+
kwargs = dict(
|
|
285
|
+
size=size,
|
|
286
|
+
fill_value=fill_value,
|
|
287
|
+
return_index=return_index,
|
|
288
|
+
return_inverse=return_inverse,
|
|
289
|
+
return_counts=return_counts,
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
if axis is not None:
|
|
293
|
+
axis_index = array._lookup_indices(axis)
|
|
294
|
+
if axis_index is None:
|
|
295
|
+
raise ValueError(f"Axis {axis} not found in array. Available axes: {array.axes}")
|
|
296
|
+
out = jnp.unique(array.array, axis=axis_index, **kwargs)
|
|
297
|
+
else:
|
|
298
|
+
out = jnp.unique(array.array, **kwargs)
|
|
299
|
+
|
|
300
|
+
if is_multireturn:
|
|
301
|
+
unique = out[0]
|
|
302
|
+
next_index = 1
|
|
303
|
+
if return_index:
|
|
304
|
+
index = out[next_index]
|
|
305
|
+
next_index += 1
|
|
306
|
+
if return_inverse:
|
|
307
|
+
inverse = out[next_index]
|
|
308
|
+
next_index += 1
|
|
309
|
+
if return_counts:
|
|
310
|
+
counts = out[next_index]
|
|
311
|
+
next_index += 1
|
|
312
|
+
else:
|
|
313
|
+
unique = out
|
|
314
|
+
|
|
315
|
+
ret = []
|
|
316
|
+
|
|
317
|
+
if axis is not None:
|
|
318
|
+
out_axes = haliax.axis.replace_axis(array.axes, axis, Unique)
|
|
319
|
+
else:
|
|
320
|
+
out_axes = (Unique,)
|
|
321
|
+
|
|
322
|
+
unique_values = haliax.named(unique, out_axes)
|
|
323
|
+
if not is_multireturn:
|
|
324
|
+
return unique_values
|
|
325
|
+
|
|
326
|
+
ret.append(unique_values)
|
|
327
|
+
|
|
328
|
+
if return_index:
|
|
329
|
+
ret.append(haliax.named(index, Unique))
|
|
330
|
+
|
|
331
|
+
if return_inverse:
|
|
332
|
+
if axis is not None:
|
|
333
|
+
assert axis_index is not None
|
|
334
|
+
inverse = haliax.named(inverse, array.axes[axis_index])
|
|
335
|
+
else:
|
|
336
|
+
inverse = haliax.named(inverse, array.axes)
|
|
337
|
+
ret.append(inverse)
|
|
338
|
+
|
|
339
|
+
if return_counts:
|
|
340
|
+
ret.append(haliax.named(counts, Unique))
|
|
341
|
+
|
|
342
|
+
return tuple(ret)
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
def unique_values(
|
|
346
|
+
array: NamedArray,
|
|
347
|
+
Unique: Axis,
|
|
348
|
+
*,
|
|
349
|
+
axis: AxisSelector | None = None,
|
|
350
|
+
fill_value: ArrayLike | None = None,
|
|
351
|
+
) -> NamedArray:
|
|
352
|
+
"""Shortcut for :func:`unique` that returns only unique values."""
|
|
353
|
+
|
|
354
|
+
return typing.cast(
|
|
355
|
+
NamedArray,
|
|
356
|
+
unique(
|
|
357
|
+
array,
|
|
358
|
+
Unique,
|
|
359
|
+
axis=axis,
|
|
360
|
+
fill_value=fill_value,
|
|
361
|
+
),
|
|
362
|
+
)
|
|
363
|
+
|
|
364
|
+
|
|
365
|
+
def unique_counts(
|
|
366
|
+
array: NamedArray,
|
|
367
|
+
Unique: Axis,
|
|
368
|
+
*,
|
|
369
|
+
axis: AxisSelector | None = None,
|
|
370
|
+
fill_value: ArrayLike | None = None,
|
|
371
|
+
) -> tuple[NamedArray, NamedArray]:
|
|
372
|
+
"""Shortcut for :func:`unique` that also returns counts."""
|
|
373
|
+
|
|
374
|
+
values, counts = typing.cast(
|
|
375
|
+
tuple[NamedArray, NamedArray],
|
|
376
|
+
unique(
|
|
377
|
+
array,
|
|
378
|
+
Unique,
|
|
379
|
+
return_counts=True,
|
|
380
|
+
axis=axis,
|
|
381
|
+
fill_value=fill_value,
|
|
382
|
+
),
|
|
383
|
+
)
|
|
384
|
+
return values, counts
|
|
385
|
+
|
|
386
|
+
|
|
387
|
+
def unique_inverse(
|
|
388
|
+
array: NamedArray,
|
|
389
|
+
Unique: Axis,
|
|
390
|
+
*,
|
|
391
|
+
axis: AxisSelector | None = None,
|
|
392
|
+
fill_value: ArrayLike | None = None,
|
|
393
|
+
) -> tuple[NamedArray, NamedArray]:
|
|
394
|
+
"""Shortcut for :func:`unique` that also returns inverse indices."""
|
|
395
|
+
|
|
396
|
+
values, inverse = typing.cast(
|
|
397
|
+
tuple[NamedArray, NamedArray],
|
|
398
|
+
unique(
|
|
399
|
+
array,
|
|
400
|
+
Unique,
|
|
401
|
+
return_inverse=True,
|
|
402
|
+
axis=axis,
|
|
403
|
+
fill_value=fill_value,
|
|
404
|
+
),
|
|
405
|
+
)
|
|
406
|
+
return values, inverse
|
|
407
|
+
|
|
408
|
+
|
|
409
|
+
def unique_all(
|
|
410
|
+
array: NamedArray,
|
|
411
|
+
Unique: Axis,
|
|
412
|
+
*,
|
|
413
|
+
axis: AxisSelector | None = None,
|
|
414
|
+
fill_value: ArrayLike | None = None,
|
|
415
|
+
) -> tuple[NamedArray, NamedArray, NamedArray, NamedArray]:
|
|
416
|
+
"""Shortcut for :func:`unique` returning values, indices, inverse, and counts."""
|
|
417
|
+
|
|
418
|
+
values, indices, inverse, counts = typing.cast(
|
|
419
|
+
tuple[NamedArray, NamedArray, NamedArray, NamedArray],
|
|
420
|
+
unique(
|
|
421
|
+
array,
|
|
422
|
+
Unique,
|
|
423
|
+
return_index=True,
|
|
424
|
+
return_inverse=True,
|
|
425
|
+
return_counts=True,
|
|
426
|
+
axis=axis,
|
|
427
|
+
fill_value=fill_value,
|
|
428
|
+
),
|
|
429
|
+
)
|
|
430
|
+
return values, indices, inverse, counts
|
|
431
|
+
|
|
432
|
+
|
|
433
|
+
__all__ = [
|
|
434
|
+
"trace",
|
|
435
|
+
"where",
|
|
436
|
+
"tril",
|
|
437
|
+
"triu",
|
|
438
|
+
"isclose",
|
|
439
|
+
"pad_left",
|
|
440
|
+
"pad",
|
|
441
|
+
"clip",
|
|
442
|
+
"unique",
|
|
443
|
+
"unique_values",
|
|
444
|
+
"unique_counts",
|
|
445
|
+
"unique_inverse",
|
|
446
|
+
"unique_all",
|
|
447
|
+
]
|
|
@@ -1,4 +1,5 @@
|
|
|
1
1
|
from typing import Callable
|
|
2
|
+
import typing
|
|
2
3
|
|
|
3
4
|
import jax.numpy as jnp
|
|
4
5
|
import pytest
|
|
@@ -250,3 +251,159 @@ def test_pad():
|
|
|
250
251
|
assert padded.axes[0].size == Height.size + 3
|
|
251
252
|
assert padded.axes[1].size == Width.size + 1
|
|
252
253
|
assert jnp.all(expected == padded.array)
|
|
254
|
+
|
|
255
|
+
|
|
256
|
+
def test_unique():
|
|
257
|
+
# named version of this test:
|
|
258
|
+
# >>> M = jnp.array([[1, 2],
|
|
259
|
+
# ... [2, 3],
|
|
260
|
+
# ... [1, 2]])
|
|
261
|
+
# >>> jnp.unique(M)
|
|
262
|
+
# Array([1, 2, 3], dtype=int32)
|
|
263
|
+
|
|
264
|
+
Height = Axis("Height", 3)
|
|
265
|
+
Width = Axis("Width", 2)
|
|
266
|
+
|
|
267
|
+
named1 = hax.named([[1, 2], [2, 3], [1, 2]], (Height, Width))
|
|
268
|
+
|
|
269
|
+
U = Axis("U", 3)
|
|
270
|
+
|
|
271
|
+
named2 = hax.unique(named1, U)
|
|
272
|
+
|
|
273
|
+
assert jnp.all(jnp.equal(named2.array, jnp.array([1, 2, 3])))
|
|
274
|
+
|
|
275
|
+
# If you pass an ``axis`` keyword, you can find unique *slices* of the array along
|
|
276
|
+
# that axis:
|
|
277
|
+
#
|
|
278
|
+
# >>> jnp.unique(M, axis=0)
|
|
279
|
+
# Array([[1, 2],
|
|
280
|
+
# [2, 3]], dtype=int32)
|
|
281
|
+
|
|
282
|
+
U2 = Axis("U2", 2)
|
|
283
|
+
named3 = hax.unique(named1, U2, axis=Height)
|
|
284
|
+
assert jnp.all(jnp.equal(named3.array, jnp.array([[1, 2], [2, 3]])))
|
|
285
|
+
|
|
286
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
287
|
+
# >>> values, indices = jnp.unique(x, return_index=True)
|
|
288
|
+
# >>> print(values)
|
|
289
|
+
# [1 3 4]
|
|
290
|
+
# >>> print(indices)
|
|
291
|
+
# [2 0 1]
|
|
292
|
+
# >>> jnp.all(values == x[indices])
|
|
293
|
+
# Array(True, dtype=bool)
|
|
294
|
+
|
|
295
|
+
x = hax.named([3, 4, 1, 3, 1], ("Height",))
|
|
296
|
+
U3 = Axis("U3", 3)
|
|
297
|
+
values, indices = hax.unique(x, U3, return_index=True)
|
|
298
|
+
|
|
299
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
300
|
+
assert jnp.all(jnp.equal(indices.array, jnp.array([2, 0, 1])))
|
|
301
|
+
|
|
302
|
+
assert jnp.all(jnp.equal(values.array, x[{"Height": indices}].array))
|
|
303
|
+
|
|
304
|
+
# If you set ``return_inverse=True``, then ``unique`` returns the indices within the
|
|
305
|
+
# unique values for every entry in the input array:
|
|
306
|
+
#
|
|
307
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
308
|
+
# >>> values, inverse = jnp.unique(x, return_inverse=True)
|
|
309
|
+
# >>> print(values)
|
|
310
|
+
# [1 3 4]
|
|
311
|
+
# >>> print(inverse)
|
|
312
|
+
# [1 2 0 1 0]
|
|
313
|
+
# >>> jnp.all(values[inverse] == x)
|
|
314
|
+
# Array(True, dtype=bool)
|
|
315
|
+
|
|
316
|
+
values, inverse = hax.unique(x, U3, return_inverse=True)
|
|
317
|
+
|
|
318
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
319
|
+
assert jnp.all(jnp.equal(inverse.array, jnp.array([1, 2, 0, 1, 0])))
|
|
320
|
+
|
|
321
|
+
# In multiple dimensions, the input can be reconstructed using
|
|
322
|
+
# :func:`jax.numpy.take`:
|
|
323
|
+
#
|
|
324
|
+
# >>> values, inverse = jnp.unique(M, axis=0, return_inverse=True)
|
|
325
|
+
# >>> jnp.all(jnp.take(values, inverse, axis=0) == M)
|
|
326
|
+
# Array(True, dtype=bool)
|
|
327
|
+
#
|
|
328
|
+
|
|
329
|
+
values, inverse = hax.unique(named1, U3, axis=Height, return_inverse=True)
|
|
330
|
+
|
|
331
|
+
assert jnp.all((values[{"U3": inverse}] == named1).array)
|
|
332
|
+
|
|
333
|
+
# **Returning counts**
|
|
334
|
+
# If you set ``return_counts=True``, then ``unique`` returns the number of occurrences
|
|
335
|
+
# within the input for every unique value:
|
|
336
|
+
#
|
|
337
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
338
|
+
# >>> values, counts = jnp.unique(x, return_counts=True)
|
|
339
|
+
# >>> print(values)
|
|
340
|
+
# [1 3 4]
|
|
341
|
+
# >>> print(counts)
|
|
342
|
+
# [2 2 1]
|
|
343
|
+
#
|
|
344
|
+
# For multi-dimensional arrays, this also returns a 1D array of counts
|
|
345
|
+
# indicating number of occurrences along the specified axis:
|
|
346
|
+
#
|
|
347
|
+
# >>> values, counts = jnp.unique(M, axis=0, return_counts=True)
|
|
348
|
+
# >>> print(values)
|
|
349
|
+
# [[1 2]
|
|
350
|
+
# [2 3]]
|
|
351
|
+
# >>> print(counts)
|
|
352
|
+
# [2 1]
|
|
353
|
+
|
|
354
|
+
values, counts = hax.unique(x, U3, return_counts=True)
|
|
355
|
+
|
|
356
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
357
|
+
assert jnp.all(jnp.equal(counts.array, jnp.array([2, 2, 1])))
|
|
358
|
+
|
|
359
|
+
values, counts = hax.unique(named1, U2, axis=Height, return_counts=True)
|
|
360
|
+
|
|
361
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([[1, 2], [2, 3]])))
|
|
362
|
+
assert jnp.all(jnp.equal(counts.array, jnp.array([2, 1])))
|
|
363
|
+
|
|
364
|
+
|
|
365
|
+
def test_unique_shortcuts():
|
|
366
|
+
Height = Axis("Height", 3)
|
|
367
|
+
Width = Axis("Width", 2)
|
|
368
|
+
|
|
369
|
+
arr2d = hax.named([[1, 2], [2, 3], [1, 2]], (Height, Width))
|
|
370
|
+
U = Axis("U", 3)
|
|
371
|
+
|
|
372
|
+
# unique_values
|
|
373
|
+
uv = hax.unique_values(arr2d, U)
|
|
374
|
+
uv_expected = hax.unique(arr2d, U)
|
|
375
|
+
assert jnp.all(uv.array == uv_expected.array)
|
|
376
|
+
|
|
377
|
+
# unique_counts
|
|
378
|
+
vc, cc = hax.unique_counts(arr2d, U)
|
|
379
|
+
vc_exp, cc_exp = hax.unique(arr2d, U, return_counts=True)
|
|
380
|
+
assert jnp.all(vc.array == vc_exp.array)
|
|
381
|
+
assert jnp.all(cc.array == cc_exp.array)
|
|
382
|
+
|
|
383
|
+
# unique_inverse
|
|
384
|
+
Height1 = Axis("Height1", 5)
|
|
385
|
+
arr1d = hax.named([3, 4, 1, 3, 1], (Height1,))
|
|
386
|
+
U2 = Axis("U2", 3)
|
|
387
|
+
vi, ii = hax.unique_inverse(arr1d, U2)
|
|
388
|
+
vi_exp, ii_exp = hax.unique(arr1d, U2, return_inverse=True)
|
|
389
|
+
assert jnp.all(vi.array == vi_exp.array)
|
|
390
|
+
assert jnp.all(ii.array == ii_exp.array)
|
|
391
|
+
|
|
392
|
+
# unique_all
|
|
393
|
+
U3 = Axis("U3", 2)
|
|
394
|
+
va, ia, ina, ca = hax.unique_all(arr2d, U3, axis=Height)
|
|
395
|
+
va_exp, ia_exp, ina_exp, ca_exp = typing.cast(
|
|
396
|
+
tuple[NamedArray, NamedArray, NamedArray, NamedArray],
|
|
397
|
+
hax.unique(
|
|
398
|
+
arr2d,
|
|
399
|
+
U3,
|
|
400
|
+
axis=Height,
|
|
401
|
+
return_index=True,
|
|
402
|
+
return_inverse=True,
|
|
403
|
+
return_counts=True,
|
|
404
|
+
),
|
|
405
|
+
)
|
|
406
|
+
assert jnp.all(va.array == va_exp.array)
|
|
407
|
+
assert jnp.all(ia.array == ia_exp.array)
|
|
408
|
+
assert jnp.all(ina.array == ina_exp.array)
|
|
409
|
+
assert jnp.all(ca.array == ca_exp.array)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev382"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|