haliax 1.4.dev381__tar.gz → 1.4.dev386__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.dev381 → haliax-1.4.dev386}/PKG-INFO +3 -3
- {haliax-1.4.dev381 → haliax-1.4.dev386}/README.md +2 -2
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/indexing.md +1 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/nn.md +1 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/state-dict.md +1 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/typing.md +1 -1
- haliax-1.4.dev386/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/__init__.py +2 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/conv.py +1 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/linear.py +1 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/ops.py +147 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_ops.py +109 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_rearrange.py +1 -1
- haliax-1.4.dev381/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.coveragerc +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.flake8 +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.gitignore +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/AGENTS.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/LICENSE +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/api.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/css/material.css +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/faq.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/fp8.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/index.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/matmul.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/partitioning.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/rearrange.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/requirements.txt +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/scan.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/tutorial.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/docs/vmap.md +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/mkdocs.yml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/pyproject.toml +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/core.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/random.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/types.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/util.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/core_test.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_attention.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_axis.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_conv.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_debug.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_dot.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_hof.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_int8.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_nn.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_pool.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_random.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_scan.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/tests/test_utils.py +0 -0
- {haliax-1.4.dev381 → haliax-1.4.dev386}/uv.lock +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev386
|
|
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/
|
|
@@ -33,14 +33,14 @@ Description-Content-Type: text/markdown
|
|
|
33
33
|
<a href="">
|
|
34
34
|
<img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
|
|
35
35
|
</a>
|
|
36
|
-
<a href="https://
|
|
36
|
+
<a href="https://pypi.org/project/haliax/">
|
|
37
37
|
<img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
|
|
38
38
|
</a>
|
|
39
39
|
|
|
40
40
|
> *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
|
|
41
41
|
> — Patrick Rothfuss, *The Name of the Wind*
|
|
42
42
|
|
|
43
|
-
Haliax is a [JAX](https
|
|
43
|
+
Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
|
|
44
44
|
Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
|
|
45
45
|
as typically used in NumPy, PyTorch, etc.
|
|
46
46
|
|
|
@@ -10,14 +10,14 @@
|
|
|
10
10
|
<a href="">
|
|
11
11
|
<img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
|
|
12
12
|
</a>
|
|
13
|
-
<a href="https://
|
|
13
|
+
<a href="https://pypi.org/project/haliax/">
|
|
14
14
|
<img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
|
|
15
15
|
</a>
|
|
16
16
|
|
|
17
17
|
> *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
|
|
18
18
|
> — Patrick Rothfuss, *The Name of the Wind*
|
|
19
19
|
|
|
20
|
-
Haliax is a [JAX](https
|
|
20
|
+
Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
|
|
21
21
|
Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
|
|
22
22
|
as typically used in NumPy, PyTorch, etc.
|
|
23
23
|
|
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
# Indexing and Slicing
|
|
2
2
|
|
|
3
3
|
Haliax supports Numpy-style indexing, including so-called [Advanced Indexing](https://numpy.org/doc/stable/user/basics.indexing.html#advanced-indexing),
|
|
4
|
-
though the syntax is necessarily different. Most forms of indexing are
|
|
4
|
+
though the syntax is necessarily different. Most forms of indexing are supported, except we don't support indexing with
|
|
5
5
|
booleans right now. (JAX doesn't support indexing with non-constant bool arrays anyway,
|
|
6
6
|
so I don't think it's worth the effort to implement it in Haliax.)
|
|
7
7
|
|
|
@@ -6,7 +6,7 @@
|
|
|
6
6
|
Haliax provides a small number of neural network modules that are compatible with Equinox, though
|
|
7
7
|
they naturally all use [haliax.NamedArray][]. (We welcome PRs for more modules! Nothing too exotic though.)
|
|
8
8
|
|
|
9
|
-
The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create
|
|
9
|
+
The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create homogeneous "stacks"
|
|
10
10
|
of the same module (e.g. transformer blocks), which is a common pattern in deep learning.
|
|
11
11
|
|
|
12
12
|
### Linear
|
|
@@ -226,7 +226,7 @@ any Axis members to match the new shape.
|
|
|
226
226
|
::: haliax.state_dict.save_state_dict
|
|
227
227
|
::: haliax.state_dict.load_state_dict
|
|
228
228
|
|
|
229
|
-
### Converting
|
|
229
|
+
### Converting between State Dicts and Modules
|
|
230
230
|
|
|
231
231
|
::: haliax.state_dict.from_state_dict
|
|
232
232
|
::: haliax.state_dict.to_state_dict
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev386"
|
|
@@ -66,7 +66,7 @@ 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 clip, isclose, pad_left, pad, trace, tril, triu, where
|
|
69
|
+
from .ops import clip, isclose, pad_left, pad, trace, tril, triu, unique, where
|
|
70
70
|
from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
|
|
71
71
|
from .specialized_fns import top_k
|
|
72
72
|
from .types import Scalar
|
|
@@ -1034,6 +1034,7 @@ __all__ = [
|
|
|
1034
1034
|
"vmap",
|
|
1035
1035
|
"trace",
|
|
1036
1036
|
"where",
|
|
1037
|
+
"unique",
|
|
1037
1038
|
"clip",
|
|
1038
1039
|
"tril",
|
|
1039
1040
|
"triu",
|
|
@@ -213,7 +213,7 @@ class Conv(_ConvBase):
|
|
|
213
213
|
return x
|
|
214
214
|
|
|
215
215
|
def _do_conv(self, inputs):
|
|
216
|
-
# _do_conv expects there
|
|
216
|
+
# _do_conv expects there to be a single __batch__ dimension
|
|
217
217
|
output_axes = _compute_output_axes(inputs, "__batch__", self.In, self.Out)
|
|
218
218
|
|
|
219
219
|
batch_index = _index_of_name(inputs.axes, "__batch__")
|
|
@@ -143,7 +143,7 @@ class MoELinear(eqx.Module):
|
|
|
143
143
|
Experts: AxisSpec = eqx.field(static=True)
|
|
144
144
|
In: Axis = eqx.field(static=True)
|
|
145
145
|
Out: Axis = eqx.field(static=True)
|
|
146
|
-
# TODO: support
|
|
146
|
+
# TODO: support quantization for ragged_dot?
|
|
147
147
|
# dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
|
|
148
148
|
|
|
149
149
|
use_gmm: bool = eqx.field(static=True)
|
|
@@ -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,147 @@ 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
|
+
__all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "pad", "clip", "unique"]
|
|
@@ -250,3 +250,112 @@ def test_pad():
|
|
|
250
250
|
assert padded.axes[0].size == Height.size + 3
|
|
251
251
|
assert padded.axes[1].size == Width.size + 1
|
|
252
252
|
assert jnp.all(expected == padded.array)
|
|
253
|
+
|
|
254
|
+
|
|
255
|
+
def test_unique():
|
|
256
|
+
# named version of this test:
|
|
257
|
+
# >>> M = jnp.array([[1, 2],
|
|
258
|
+
# ... [2, 3],
|
|
259
|
+
# ... [1, 2]])
|
|
260
|
+
# >>> jnp.unique(M)
|
|
261
|
+
# Array([1, 2, 3], dtype=int32)
|
|
262
|
+
|
|
263
|
+
Height = Axis("Height", 3)
|
|
264
|
+
Width = Axis("Width", 2)
|
|
265
|
+
|
|
266
|
+
named1 = hax.named([[1, 2], [2, 3], [1, 2]], (Height, Width))
|
|
267
|
+
|
|
268
|
+
U = Axis("U", 3)
|
|
269
|
+
|
|
270
|
+
named2 = hax.unique(named1, U)
|
|
271
|
+
|
|
272
|
+
assert jnp.all(jnp.equal(named2.array, jnp.array([1, 2, 3])))
|
|
273
|
+
|
|
274
|
+
# If you pass an ``axis`` keyword, you can find unique *slices* of the array along
|
|
275
|
+
# that axis:
|
|
276
|
+
#
|
|
277
|
+
# >>> jnp.unique(M, axis=0)
|
|
278
|
+
# Array([[1, 2],
|
|
279
|
+
# [2, 3]], dtype=int32)
|
|
280
|
+
|
|
281
|
+
U2 = Axis("U2", 2)
|
|
282
|
+
named3 = hax.unique(named1, U2, axis=Height)
|
|
283
|
+
assert jnp.all(jnp.equal(named3.array, jnp.array([[1, 2], [2, 3]])))
|
|
284
|
+
|
|
285
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
286
|
+
# >>> values, indices = jnp.unique(x, return_index=True)
|
|
287
|
+
# >>> print(values)
|
|
288
|
+
# [1 3 4]
|
|
289
|
+
# >>> print(indices)
|
|
290
|
+
# [2 0 1]
|
|
291
|
+
# >>> jnp.all(values == x[indices])
|
|
292
|
+
# Array(True, dtype=bool)
|
|
293
|
+
|
|
294
|
+
x = hax.named([3, 4, 1, 3, 1], ("Height",))
|
|
295
|
+
U3 = Axis("U3", 3)
|
|
296
|
+
values, indices = hax.unique(x, U3, return_index=True)
|
|
297
|
+
|
|
298
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
299
|
+
assert jnp.all(jnp.equal(indices.array, jnp.array([2, 0, 1])))
|
|
300
|
+
|
|
301
|
+
assert jnp.all(jnp.equal(values.array, x[{"Height": indices}].array))
|
|
302
|
+
|
|
303
|
+
# If you set ``return_inverse=True``, then ``unique`` returns the indices within the
|
|
304
|
+
# unique values for every entry in the input array:
|
|
305
|
+
#
|
|
306
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
307
|
+
# >>> values, inverse = jnp.unique(x, return_inverse=True)
|
|
308
|
+
# >>> print(values)
|
|
309
|
+
# [1 3 4]
|
|
310
|
+
# >>> print(inverse)
|
|
311
|
+
# [1 2 0 1 0]
|
|
312
|
+
# >>> jnp.all(values[inverse] == x)
|
|
313
|
+
# Array(True, dtype=bool)
|
|
314
|
+
|
|
315
|
+
values, inverse = hax.unique(x, U3, return_inverse=True)
|
|
316
|
+
|
|
317
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
318
|
+
assert jnp.all(jnp.equal(inverse.array, jnp.array([1, 2, 0, 1, 0])))
|
|
319
|
+
|
|
320
|
+
# In multiple dimensions, the input can be reconstructed using
|
|
321
|
+
# :func:`jax.numpy.take`:
|
|
322
|
+
#
|
|
323
|
+
# >>> values, inverse = jnp.unique(M, axis=0, return_inverse=True)
|
|
324
|
+
# >>> jnp.all(jnp.take(values, inverse, axis=0) == M)
|
|
325
|
+
# Array(True, dtype=bool)
|
|
326
|
+
#
|
|
327
|
+
|
|
328
|
+
values, inverse = hax.unique(named1, U3, axis=Height, return_inverse=True)
|
|
329
|
+
|
|
330
|
+
assert jnp.all((values[{"U3": inverse}] == named1).array)
|
|
331
|
+
|
|
332
|
+
# **Returning counts**
|
|
333
|
+
# If you set ``return_counts=True``, then ``unique`` returns the number of occurrences
|
|
334
|
+
# within the input for every unique value:
|
|
335
|
+
#
|
|
336
|
+
# >>> x = jnp.array([3, 4, 1, 3, 1])
|
|
337
|
+
# >>> values, counts = jnp.unique(x, return_counts=True)
|
|
338
|
+
# >>> print(values)
|
|
339
|
+
# [1 3 4]
|
|
340
|
+
# >>> print(counts)
|
|
341
|
+
# [2 2 1]
|
|
342
|
+
#
|
|
343
|
+
# For multi-dimensional arrays, this also returns a 1D array of counts
|
|
344
|
+
# indicating number of occurrences along the specified axis:
|
|
345
|
+
#
|
|
346
|
+
# >>> values, counts = jnp.unique(M, axis=0, return_counts=True)
|
|
347
|
+
# >>> print(values)
|
|
348
|
+
# [[1 2]
|
|
349
|
+
# [2 3]]
|
|
350
|
+
# >>> print(counts)
|
|
351
|
+
# [2 1]
|
|
352
|
+
|
|
353
|
+
values, counts = hax.unique(x, U3, return_counts=True)
|
|
354
|
+
|
|
355
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
|
|
356
|
+
assert jnp.all(jnp.equal(counts.array, jnp.array([2, 2, 1])))
|
|
357
|
+
|
|
358
|
+
values, counts = hax.unique(named1, U2, axis=Height, return_counts=True)
|
|
359
|
+
|
|
360
|
+
assert jnp.all(jnp.equal(values.array, jnp.array([[1, 2], [2, 3]])))
|
|
361
|
+
assert jnp.all(jnp.equal(counts.array, jnp.array([2, 1])))
|
|
@@ -293,7 +293,7 @@ def test_examples():
|
|
|
293
293
|
r = einops_rearrange(z, "{B (H: h1 h) (W: w1 w) C} -> (B: B h1 w1) ... (C: C h w) ", h1=2, w1=2)
|
|
294
294
|
assert r.axes == (Axis("B", B.size * 2 * 2), D, Axis("C", C.size * sH.size * sW.size))
|
|
295
295
|
# unet attention reordering:
|
|
296
|
-
#
|
|
296
|
+
# positional: (qkv heads c) h w -> qkv heads c (h w)
|
|
297
297
|
# named: { (embed: qkv heads c) h w } -> qkv heads c (pos: h w)
|
|
298
298
|
Embed = Axis("embed", 3 * 4 * C.size)
|
|
299
299
|
attn = hax.random.randint(PRNGKey(0), (Embed, H, W), 0, 255)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev381"
|
|
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
|