haliax 1.4.dev407__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.dev407 → haliax-1.4.dev408}/.agents/projects/api_parity.md +12 -12
- {haliax-1.4.dev407 → haliax-1.4.dev408}/PKG-INFO +1 -1
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/api.md +12 -0
- haliax-1.4.dev408/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/__init__.py +111 -0
- haliax-1.4.dev408/tests/test_nan_reductions.py +59 -0
- haliax-1.4.dev407/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.coveragerc +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.flake8 +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.gitignore +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/AGENTS.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/LICENSE +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/README.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/css/material.css +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/faq.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/fp8.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/index.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/indexing.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/matmul.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/nn.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/partitioning.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/primer.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/rearrange.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/requirements.txt +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/scan.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/state-dict.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/tutorial.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/typing.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/docs/vmap.md +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/mkdocs.yml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/pyproject.toml +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/core.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/field.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/random.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/types.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/util.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/core_test.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_attention.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_axis.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_conv.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_debug.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_dot.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_field.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_hof.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_int8.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_nn.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_ops.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_pool.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_random.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_scan.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_utils.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev407 → haliax-1.4.dev408}/uv.lock +0 -0
|
@@ -5,7 +5,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
|
|
|
5
5
|
|
|
6
6
|
## numpy
|
|
7
7
|
- [x] `allclose`
|
|
8
|
-
- [
|
|
8
|
+
- [x] `amin`
|
|
9
9
|
- [ ] `append`
|
|
10
10
|
- [ ] `apply_along_axis`
|
|
11
11
|
- [ ] `apply_over_axes`
|
|
@@ -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
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev408"
|
|
@@ -560,6 +560,13 @@ def amax(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Opti
|
|
|
560
560
|
return wrap_reduction_call(jnp.amax, array, axis, where, single_axis_only=False, supports_where=True)
|
|
561
561
|
|
|
562
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
|
+
|
|
563
570
|
def any(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
|
|
564
571
|
"""True if any elements along a given axis or axes are True. If axis is None, any elements are True."""
|
|
565
572
|
return wrap_reduction_call(jnp.any, array, axis, where, single_axis_only=False, supports_where=True)
|
|
@@ -656,6 +663,84 @@ def var(
|
|
|
656
663
|
)
|
|
657
664
|
|
|
658
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
|
+
|
|
659
744
|
# "Normalization" functions that use an axis but don't change the shape
|
|
660
745
|
|
|
661
746
|
|
|
@@ -673,6 +758,20 @@ def cumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None
|
|
|
673
758
|
return wrap_axiswise_call(jnp.cumprod, a, axis, dtype=dtype, single_axis_only=True)
|
|
674
759
|
|
|
675
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
|
+
|
|
676
775
|
def sort(a: NamedArray, axis: AxisSelector) -> NamedArray:
|
|
677
776
|
"""
|
|
678
777
|
Named version of [jax.numpy.sort](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.sort.html)
|
|
@@ -1034,12 +1133,22 @@ __all__ = [
|
|
|
1034
1133
|
"trunc",
|
|
1035
1134
|
"all",
|
|
1036
1135
|
"amax",
|
|
1136
|
+
"amin",
|
|
1037
1137
|
"any",
|
|
1038
1138
|
"argmax",
|
|
1039
1139
|
"argmin",
|
|
1040
1140
|
"max",
|
|
1041
1141
|
"mean",
|
|
1042
1142
|
"min",
|
|
1143
|
+
"nanargmax",
|
|
1144
|
+
"nanargmin",
|
|
1145
|
+
"nanmax",
|
|
1146
|
+
"nanmean",
|
|
1147
|
+
"nanmin",
|
|
1148
|
+
"nanprod",
|
|
1149
|
+
"nanstd",
|
|
1150
|
+
"nansum",
|
|
1151
|
+
"nanvar",
|
|
1043
1152
|
"prod",
|
|
1044
1153
|
"product",
|
|
1045
1154
|
"ptp",
|
|
@@ -1048,6 +1157,8 @@ __all__ = [
|
|
|
1048
1157
|
"var",
|
|
1049
1158
|
"cumsum",
|
|
1050
1159
|
"cumprod",
|
|
1160
|
+
"nancumprod",
|
|
1161
|
+
"nancumsum",
|
|
1051
1162
|
"sort",
|
|
1052
1163
|
"scan",
|
|
1053
1164
|
"fold",
|
|
@@ -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)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev407"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|