haliax 1.4.dev406__tar.gz → 1.4.dev408__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.dev406 → haliax-1.4.dev408}/.agents/projects/api_parity.md +15 -15
- {haliax-1.4.dev406 → haliax-1.4.dev408}/PKG-INFO +1 -1
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/api.md +15 -0
- haliax-1.4.dev408/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/__init__.py +117 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/ops.py +23 -0
- haliax-1.4.dev408/tests/test_nan_reductions.py +59 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_ops.py +26 -0
- haliax-1.4.dev406/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.coveragerc +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.flake8 +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.gitignore +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/AGENTS.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/LICENSE +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/README.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/css/material.css +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/faq.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/fp8.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/index.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/indexing.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/matmul.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/nn.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/partitioning.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/primer.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/rearrange.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/requirements.txt +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/scan.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/state-dict.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/tutorial.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/typing.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/vmap.md +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/mkdocs.yml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/pyproject.toml +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/core.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/field.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/random.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/types.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/util.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/core_test.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_attention.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_axis.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_conv.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_debug.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_dot.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_field.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_hof.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_int8.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_nn.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_pool.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_random.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_scan.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_utils.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev406 → haliax-1.4.dev408}/uv.lock +0 -0
|
@@ -4,15 +4,15 @@ This document tracks JAX NumPy functions not yet wrapped by Haliax.
|
|
|
4
4
|
APIs that don't translate well to named tensors are intentionally omitted here. This includes dtype constructors, raw array converters (e.g. `from_dlpack`), indexing helpers like `c_`/`r_`, and functions whose JAX counterparts already work with `NamedArray` out of the box. Basic array construction is handled by `haliax.named`, so functions such as `array` or `asarray` are not listed.
|
|
5
5
|
|
|
6
6
|
## numpy
|
|
7
|
-
- [
|
|
8
|
-
- [
|
|
7
|
+
- [x] `allclose`
|
|
8
|
+
- [x] `amin`
|
|
9
9
|
- [ ] `append`
|
|
10
10
|
- [ ] `apply_along_axis`
|
|
11
11
|
- [ ] `apply_over_axes`
|
|
12
12
|
- [ ] `argpartition`
|
|
13
13
|
- [ ] `argwhere`
|
|
14
|
-
- [
|
|
15
|
-
- [
|
|
14
|
+
- [x] `array_equal`
|
|
15
|
+
- [x] `array_equiv`
|
|
16
16
|
- [ ] `array_split`
|
|
17
17
|
- [ ] `astype`
|
|
18
18
|
- [ ] `atan2`
|
|
@@ -93,20 +93,20 @@ APIs that don't translate well to named tensors are intentionally omitted here.
|
|
|
93
93
|
- [ ] `modf`
|
|
94
94
|
- [ ] `moveaxis`
|
|
95
95
|
- [ ] `nan_to_num`
|
|
96
|
-
- [
|
|
97
|
-
- [
|
|
98
|
-
- [
|
|
99
|
-
- [
|
|
100
|
-
- [
|
|
101
|
-
- [
|
|
96
|
+
- [x] `nanargmax`
|
|
97
|
+
- [x] `nanargmin`
|
|
98
|
+
- [x] `nancumprod`
|
|
99
|
+
- [x] `nancumsum`
|
|
100
|
+
- [x] `nanmax`
|
|
101
|
+
- [x] `nanmean`
|
|
102
102
|
- [ ] `nanmedian`
|
|
103
|
-
- [
|
|
103
|
+
- [x] `nanmin`
|
|
104
104
|
- [ ] `nanpercentile`
|
|
105
|
-
- [
|
|
105
|
+
- [x] `nanprod`
|
|
106
106
|
- [ ] `nanquantile`
|
|
107
|
-
- [
|
|
108
|
-
- [
|
|
109
|
-
- [
|
|
107
|
+
- [x] `nanstd`
|
|
108
|
+
- [x] `nansum`
|
|
109
|
+
- [x] `nanvar`
|
|
110
110
|
- [ ] `nonzero`
|
|
111
111
|
- [ ] `ogrid`
|
|
112
112
|
- [ ] `packbits`
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev408
|
|
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/
|
|
@@ -128,12 +128,22 @@ You can convert it to a [jax.numpy.ndarray][] with [haliax.NamedArray.scalar][],
|
|
|
128
128
|
|
|
129
129
|
::: haliax.all
|
|
130
130
|
::: haliax.amax
|
|
131
|
+
::: haliax.amin
|
|
131
132
|
::: haliax.any
|
|
132
133
|
::: haliax.argmax
|
|
133
134
|
::: haliax.argmin
|
|
134
135
|
::: haliax.max
|
|
135
136
|
::: haliax.mean
|
|
136
137
|
::: haliax.min
|
|
138
|
+
::: haliax.nanargmax
|
|
139
|
+
::: haliax.nanargmin
|
|
140
|
+
::: haliax.nanmax
|
|
141
|
+
::: haliax.nanmean
|
|
142
|
+
::: haliax.nanmin
|
|
143
|
+
::: haliax.nanprod
|
|
144
|
+
::: haliax.nanstd
|
|
145
|
+
::: haliax.nansum
|
|
146
|
+
::: haliax.nanvar
|
|
137
147
|
::: haliax.prod
|
|
138
148
|
::: haliax.ptp
|
|
139
149
|
::: haliax.std
|
|
@@ -146,6 +156,8 @@ don't reduce it.
|
|
|
146
156
|
|
|
147
157
|
::: haliax.cumsum
|
|
148
158
|
::: haliax.cumprod
|
|
159
|
+
::: haliax.nancumprod
|
|
160
|
+
::: haliax.nancumsum
|
|
149
161
|
::: haliax.sort
|
|
150
162
|
::: haliax.argsort
|
|
151
163
|
|
|
@@ -259,6 +271,9 @@ These are all more or less directly from JAX's NumPy API.
|
|
|
259
271
|
::: haliax.bincount
|
|
260
272
|
::: haliax.clip
|
|
261
273
|
::: haliax.isclose
|
|
274
|
+
::: haliax.allclose
|
|
275
|
+
::: haliax.array_equal
|
|
276
|
+
::: haliax.array_equiv
|
|
262
277
|
::: haliax.pad
|
|
263
278
|
::: haliax.searchsorted
|
|
264
279
|
::: haliax.top_k
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev408"
|
|
@@ -69,6 +69,9 @@ from .hof import fold, map, scan, vmap
|
|
|
69
69
|
from .jax_utils import tree_checkpoint_name
|
|
70
70
|
from .ops import (
|
|
71
71
|
clip,
|
|
72
|
+
allclose,
|
|
73
|
+
array_equal,
|
|
74
|
+
array_equiv,
|
|
72
75
|
isclose,
|
|
73
76
|
pad_left,
|
|
74
77
|
pad,
|
|
@@ -557,6 +560,13 @@ def amax(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Opti
|
|
|
557
560
|
return wrap_reduction_call(jnp.amax, array, axis, where, single_axis_only=False, supports_where=True)
|
|
558
561
|
|
|
559
562
|
|
|
563
|
+
def amin(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
|
|
564
|
+
"""
|
|
565
|
+
Aliax for min. See min for details.
|
|
566
|
+
"""
|
|
567
|
+
return wrap_reduction_call(jnp.amin, array, axis, where, single_axis_only=False, supports_where=True)
|
|
568
|
+
|
|
569
|
+
|
|
560
570
|
def any(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
|
|
561
571
|
"""True if any elements along a given axis or axes are True. If axis is None, any elements are True."""
|
|
562
572
|
return wrap_reduction_call(jnp.any, array, axis, where, single_axis_only=False, supports_where=True)
|
|
@@ -653,6 +663,84 @@ def var(
|
|
|
653
663
|
)
|
|
654
664
|
|
|
655
665
|
|
|
666
|
+
def nanargmax(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
|
|
667
|
+
return wrap_reduction_call(jnp.nanargmax, array, axis, None, single_axis_only=True, supports_where=False)
|
|
668
|
+
|
|
669
|
+
|
|
670
|
+
def nanargmin(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
|
|
671
|
+
return wrap_reduction_call(jnp.nanargmin, array, axis, None, single_axis_only=True, supports_where=False)
|
|
672
|
+
|
|
673
|
+
|
|
674
|
+
def nanmax(
|
|
675
|
+
array: NamedArray,
|
|
676
|
+
axis: Optional[AxisSelection] = None,
|
|
677
|
+
*,
|
|
678
|
+
where: Optional[NamedArray] = None,
|
|
679
|
+
) -> NamedArray:
|
|
680
|
+
return wrap_reduction_call(jnp.nanmax, array, axis, where, single_axis_only=False, supports_where=True)
|
|
681
|
+
|
|
682
|
+
|
|
683
|
+
def nanmean(
|
|
684
|
+
array: NamedArray,
|
|
685
|
+
axis: Optional[AxisSelection] = None,
|
|
686
|
+
*,
|
|
687
|
+
where: Optional[NamedArray] = None,
|
|
688
|
+
dtype: Optional[DTypeLike] = None,
|
|
689
|
+
) -> NamedArray:
|
|
690
|
+
return wrap_reduction_call(jnp.nanmean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
|
|
691
|
+
|
|
692
|
+
|
|
693
|
+
def nanmin(
|
|
694
|
+
array: NamedArray,
|
|
695
|
+
axis: Optional[AxisSelection] = None,
|
|
696
|
+
*,
|
|
697
|
+
where: Optional[NamedArray] = None,
|
|
698
|
+
) -> NamedArray:
|
|
699
|
+
return wrap_reduction_call(jnp.nanmin, array, axis, where, single_axis_only=False, supports_where=True)
|
|
700
|
+
|
|
701
|
+
|
|
702
|
+
def nanprod(
|
|
703
|
+
array: NamedArray,
|
|
704
|
+
axis: Optional[AxisSelection] = None,
|
|
705
|
+
*,
|
|
706
|
+
where: Optional[NamedArray] = None,
|
|
707
|
+
dtype: Optional[DTypeLike] = None,
|
|
708
|
+
) -> NamedArray:
|
|
709
|
+
return wrap_reduction_call(jnp.nanprod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
|
|
710
|
+
|
|
711
|
+
|
|
712
|
+
def nanstd(
|
|
713
|
+
array: NamedArray,
|
|
714
|
+
axis: Optional[AxisSelection] = None,
|
|
715
|
+
*,
|
|
716
|
+
where: Optional[NamedArray] = None,
|
|
717
|
+
ddof: int = 0,
|
|
718
|
+
dtype: Optional[DTypeLike] = None,
|
|
719
|
+
) -> NamedArray:
|
|
720
|
+
return wrap_reduction_call(jnp.nanstd, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof)
|
|
721
|
+
|
|
722
|
+
|
|
723
|
+
def nansum(
|
|
724
|
+
array: NamedArray,
|
|
725
|
+
axis: Optional[AxisSelection] = None,
|
|
726
|
+
*,
|
|
727
|
+
where: Optional[NamedArray] = None,
|
|
728
|
+
dtype: Optional[DTypeLike] = None,
|
|
729
|
+
) -> NamedArray:
|
|
730
|
+
return wrap_reduction_call(jnp.nansum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
|
|
731
|
+
|
|
732
|
+
|
|
733
|
+
def nanvar(
|
|
734
|
+
array: NamedArray,
|
|
735
|
+
axis: Optional[AxisSelection] = None,
|
|
736
|
+
*,
|
|
737
|
+
where: Optional[NamedArray] = None,
|
|
738
|
+
ddof: int = 0,
|
|
739
|
+
dtype: Optional[DTypeLike] = None,
|
|
740
|
+
) -> NamedArray:
|
|
741
|
+
return wrap_reduction_call(jnp.nanvar, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof)
|
|
742
|
+
|
|
743
|
+
|
|
656
744
|
# "Normalization" functions that use an axis but don't change the shape
|
|
657
745
|
|
|
658
746
|
|
|
@@ -670,6 +758,20 @@ def cumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None
|
|
|
670
758
|
return wrap_axiswise_call(jnp.cumprod, a, axis, dtype=dtype, single_axis_only=True)
|
|
671
759
|
|
|
672
760
|
|
|
761
|
+
def nancumsum(a: NamedArray, axis: AxisSelector, *, dtype: Optional[DTypeLike] = None) -> NamedArray:
|
|
762
|
+
"""
|
|
763
|
+
Named version of [jax.numpy.nancumsum](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumsum.html)
|
|
764
|
+
"""
|
|
765
|
+
return wrap_axiswise_call(jnp.nancumsum, a, axis, dtype=dtype, single_axis_only=True)
|
|
766
|
+
|
|
767
|
+
|
|
768
|
+
def nancumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None) -> NamedArray:
|
|
769
|
+
"""
|
|
770
|
+
Named version of [jax.numpy.nancumprod](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumprod.html)
|
|
771
|
+
"""
|
|
772
|
+
return wrap_axiswise_call(jnp.nancumprod, a, axis, dtype=dtype, single_axis_only=True)
|
|
773
|
+
|
|
774
|
+
|
|
673
775
|
def sort(a: NamedArray, axis: AxisSelector) -> NamedArray:
|
|
674
776
|
"""
|
|
675
777
|
Named version of [jax.numpy.sort](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.sort.html)
|
|
@@ -1031,12 +1133,22 @@ __all__ = [
|
|
|
1031
1133
|
"trunc",
|
|
1032
1134
|
"all",
|
|
1033
1135
|
"amax",
|
|
1136
|
+
"amin",
|
|
1034
1137
|
"any",
|
|
1035
1138
|
"argmax",
|
|
1036
1139
|
"argmin",
|
|
1037
1140
|
"max",
|
|
1038
1141
|
"mean",
|
|
1039
1142
|
"min",
|
|
1143
|
+
"nanargmax",
|
|
1144
|
+
"nanargmin",
|
|
1145
|
+
"nanmax",
|
|
1146
|
+
"nanmean",
|
|
1147
|
+
"nanmin",
|
|
1148
|
+
"nanprod",
|
|
1149
|
+
"nanstd",
|
|
1150
|
+
"nansum",
|
|
1151
|
+
"nanvar",
|
|
1040
1152
|
"prod",
|
|
1041
1153
|
"product",
|
|
1042
1154
|
"ptp",
|
|
@@ -1045,6 +1157,8 @@ __all__ = [
|
|
|
1045
1157
|
"var",
|
|
1046
1158
|
"cumsum",
|
|
1047
1159
|
"cumprod",
|
|
1160
|
+
"nancumprod",
|
|
1161
|
+
"nancumsum",
|
|
1048
1162
|
"sort",
|
|
1049
1163
|
"scan",
|
|
1050
1164
|
"fold",
|
|
@@ -1105,6 +1219,9 @@ __all__ = [
|
|
|
1105
1219
|
"shard",
|
|
1106
1220
|
"enable_shape_checks",
|
|
1107
1221
|
"are_shape_checks_enabled",
|
|
1222
|
+
"allclose",
|
|
1223
|
+
"array_equal",
|
|
1224
|
+
"array_equiv",
|
|
1108
1225
|
"isclose",
|
|
1109
1226
|
"pad_left",
|
|
1110
1227
|
"pad",
|
|
@@ -138,6 +138,29 @@ def isclose(a: NamedArray, b: NamedArray, rtol=1e-05, atol=1e-08, equal_nan=Fals
|
|
|
138
138
|
return NamedArray(jnp.isclose(a.array, b.array, rtol=rtol, atol=atol, equal_nan=equal_nan), a.axes)
|
|
139
139
|
|
|
140
140
|
|
|
141
|
+
def allclose(a: NamedArray, b: NamedArray, rtol=1e-05, atol=1e-08, equal_nan=False) -> bool:
|
|
142
|
+
"""Returns True if two arrays are element-wise equal within a tolerance."""
|
|
143
|
+
a, b = broadcast_arrays(a, b)
|
|
144
|
+
return bool(jnp.allclose(a.array, b.array, rtol=rtol, atol=atol, equal_nan=equal_nan))
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def array_equal(a: NamedArray, b: NamedArray) -> bool:
|
|
148
|
+
"""Returns True if two arrays have the same shape and elements."""
|
|
149
|
+
if set(a.axes) != set(b.axes):
|
|
150
|
+
return False
|
|
151
|
+
b = b.rearrange(a.axes)
|
|
152
|
+
return bool(jnp.array_equal(a.array, b.array))
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
def array_equiv(a: NamedArray, b: NamedArray) -> bool:
|
|
156
|
+
"""Returns True if two arrays are shape-consistent and equal."""
|
|
157
|
+
try:
|
|
158
|
+
a, b = broadcast_arrays(a, b)
|
|
159
|
+
except ValueError:
|
|
160
|
+
return False
|
|
161
|
+
return bool(jnp.array_equal(a.array, b.array))
|
|
162
|
+
|
|
163
|
+
|
|
141
164
|
def pad_left(array: NamedArray, axis: Axis, new_axis: Axis, value=0) -> NamedArray:
|
|
142
165
|
"""Pad an array along named axes."""
|
|
143
166
|
amount_to_pad_to = new_axis.size - axis.size
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
from typing import Any, Callable
|
|
2
|
+
|
|
3
|
+
import jax.numpy as jnp
|
|
4
|
+
import haliax as hax
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def _sample_array():
|
|
8
|
+
Height, Width = hax.make_axes(Height=2, Width=3)
|
|
9
|
+
data = jnp.array([[1.0, jnp.nan, 3.0], [jnp.nan, 5.0, 6.0]])
|
|
10
|
+
arr = hax.named(data, (Height, Width))
|
|
11
|
+
return Height, Width, data, arr
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def test_amin_alias():
|
|
15
|
+
Height, Width = hax.make_axes(Height=2, Width=3)
|
|
16
|
+
data = jnp.arange(6.0).reshape(2, 3)
|
|
17
|
+
arr = hax.named(data, (Height, Width))
|
|
18
|
+
assert jnp.array_equal(hax.amin(arr).array, jnp.amin(data))
|
|
19
|
+
assert jnp.array_equal(hax.amin(arr, axis=Height).array, jnp.amin(data, axis=0))
|
|
20
|
+
assert hax.amin(arr, axis=Height).axes == (Width,)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def test_nan_reductions():
|
|
24
|
+
Height, Width, data, arr = _sample_array()
|
|
25
|
+
|
|
26
|
+
funcs: list[tuple[Callable[..., Any], Callable[..., Any]]] = [
|
|
27
|
+
(hax.nanmin, jnp.nanmin),
|
|
28
|
+
(hax.nanmax, jnp.nanmax),
|
|
29
|
+
(hax.nanmean, jnp.nanmean),
|
|
30
|
+
(hax.nansum, jnp.nansum),
|
|
31
|
+
(hax.nanprod, jnp.nanprod),
|
|
32
|
+
(hax.nanstd, jnp.nanstd),
|
|
33
|
+
(hax.nanvar, jnp.nanvar),
|
|
34
|
+
]
|
|
35
|
+
|
|
36
|
+
for hfunc, jfunc in funcs:
|
|
37
|
+
assert jnp.allclose(hfunc(arr).array, jfunc(data), equal_nan=True)
|
|
38
|
+
assert jnp.allclose(hfunc(arr, axis=Height).array, jfunc(data, axis=0), equal_nan=True)
|
|
39
|
+
assert hfunc(arr, axis=Height).axes == (Width,)
|
|
40
|
+
|
|
41
|
+
arg_funcs: list[tuple[Callable[..., Any], Callable[..., Any]]] = [
|
|
42
|
+
(hax.nanargmax, jnp.nanargmax),
|
|
43
|
+
(hax.nanargmin, jnp.nanargmin),
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
for hfunc, jfunc in arg_funcs:
|
|
47
|
+
out = hfunc(arr, axis=Height)
|
|
48
|
+
assert jnp.array_equal(out.array, jfunc(data, axis=0))
|
|
49
|
+
assert out.axes == (Width,)
|
|
50
|
+
|
|
51
|
+
axiswise_funcs: list[tuple[Callable[..., Any], Callable[..., Any]]] = [
|
|
52
|
+
(hax.nancumsum, jnp.nancumsum),
|
|
53
|
+
(hax.nancumprod, jnp.nancumprod),
|
|
54
|
+
]
|
|
55
|
+
|
|
56
|
+
for hfunc, jfunc in axiswise_funcs:
|
|
57
|
+
out = hfunc(arr, axis=Height)
|
|
58
|
+
assert jnp.allclose(out.array, jfunc(data, axis=0), equal_nan=True)
|
|
59
|
+
assert out.axes == (Height, Width)
|
|
@@ -425,6 +425,32 @@ def test_bincount():
|
|
|
425
425
|
assert jnp.allclose(out_w.array, expected_w)
|
|
426
426
|
|
|
427
427
|
|
|
428
|
+
def test_allclose_array_equal_equiv():
|
|
429
|
+
A = Axis("A", 2)
|
|
430
|
+
B = Axis("B", 3)
|
|
431
|
+
x = hax.random.uniform(PRNGKey(0), (A, B))
|
|
432
|
+
y = x + 1e-6
|
|
433
|
+
|
|
434
|
+
assert hax.allclose(x, y)
|
|
435
|
+
assert not hax.allclose(x, x + 1.0)
|
|
436
|
+
|
|
437
|
+
x1 = hax.ones((A, B))
|
|
438
|
+
y_reordered = x1.rearrange((B, A))
|
|
439
|
+
assert hax.array_equal(x1, y_reordered)
|
|
440
|
+
|
|
441
|
+
scalar = hax.ones(())
|
|
442
|
+
assert hax.array_equiv(x1, scalar)
|
|
443
|
+
assert not hax.array_equal(x1, scalar)
|
|
444
|
+
|
|
445
|
+
y_vec = hax.ones((B,))
|
|
446
|
+
assert hax.array_equiv(x1, y_vec)
|
|
447
|
+
assert not hax.array_equal(x1, y_vec)
|
|
448
|
+
|
|
449
|
+
C = Axis("C", 4)
|
|
450
|
+
z = hax.ones((C,))
|
|
451
|
+
assert not hax.array_equiv(x1, z)
|
|
452
|
+
|
|
453
|
+
|
|
428
454
|
def test_roll_scalar_named_shift():
|
|
429
455
|
H = Axis("H", 4)
|
|
430
456
|
W = Axis("W", 3)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev406"
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|