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.
Files changed (122) hide show
  1. {haliax-1.4.dev406 → haliax-1.4.dev408}/.agents/projects/api_parity.md +15 -15
  2. {haliax-1.4.dev406 → haliax-1.4.dev408}/PKG-INFO +1 -1
  3. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/api.md +15 -0
  4. haliax-1.4.dev408/src/haliax/__about__.py +1 -0
  5. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/__init__.py +117 -0
  6. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/ops.py +23 -0
  7. haliax-1.4.dev408/tests/test_nan_reductions.py +59 -0
  8. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_ops.py +26 -0
  9. haliax-1.4.dev406/src/haliax/__about__.py +0 -1
  10. {haliax-1.4.dev406 → haliax-1.4.dev408}/.coveragerc +0 -0
  11. {haliax-1.4.dev406 → haliax-1.4.dev408}/.flake8 +0 -0
  12. {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/publish_dev.yaml +0 -0
  13. {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_pre_commit.yaml +0 -0
  14. {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  15. {haliax-1.4.dev406 → haliax-1.4.dev408}/.github/workflows/run_tests.yaml +0 -0
  16. {haliax-1.4.dev406 → haliax-1.4.dev408}/.gitignore +0 -0
  17. {haliax-1.4.dev406 → haliax-1.4.dev408}/.playbooks/add-types.md +0 -0
  18. {haliax-1.4.dev406 → haliax-1.4.dev408}/.playbooks/wrap-non-named.md +0 -0
  19. {haliax-1.4.dev406 → haliax-1.4.dev408}/.pre-commit-config.yaml +0 -0
  20. {haliax-1.4.dev406 → haliax-1.4.dev408}/.readthedocs.yaml +0 -0
  21. {haliax-1.4.dev406 → haliax-1.4.dev408}/AGENTS.md +0 -0
  22. {haliax-1.4.dev406 → haliax-1.4.dev408}/CONTRIBUTING.md +0 -0
  23. {haliax-1.4.dev406 → haliax-1.4.dev408}/LICENSE +0 -0
  24. {haliax-1.4.dev406 → haliax-1.4.dev408}/README.md +0 -0
  25. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/broadcasting.md +0 -0
  26. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/cheatsheet.md +0 -0
  27. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/css/material.css +0 -0
  28. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/css/mkdocstrings.css +0 -0
  29. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/faq.md +0 -0
  30. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh.png +0 -0
  31. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  32. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_1d.png +0 -0
  33. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_1d_zero.png +0 -0
  34. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d.png +0 -0
  35. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  36. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  37. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  38. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  39. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/figures/device_mesh_2d_zero.png +0 -0
  40. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/fp8.md +0 -0
  41. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/index.md +0 -0
  42. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/indexing.md +0 -0
  43. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/matmul.md +0 -0
  44. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/nn.md +0 -0
  45. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/partitioning.md +0 -0
  46. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/primer.md +0 -0
  47. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/rearrange.ipynb +0 -0
  48. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/rearrange.md +0 -0
  49. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/requirements.txt +0 -0
  50. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/scan.md +0 -0
  51. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/state-dict.md +0 -0
  52. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/tutorial.md +0 -0
  53. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/typing.md +0 -0
  54. {haliax-1.4.dev406 → haliax-1.4.dev408}/docs/vmap.md +0 -0
  55. {haliax-1.4.dev406 → haliax-1.4.dev408}/mkdocs.yml +0 -0
  56. {haliax-1.4.dev406 → haliax-1.4.dev408}/pyproject.toml +0 -0
  57. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/__init__.py +0 -0
  58. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/compile_utils.py +0 -0
  59. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/dot.py +0 -0
  60. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/einsum.py +0 -0
  61. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/fp8.py +0 -0
  62. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/parsing.py +0 -0
  63. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/rearrange.py +0 -0
  64. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/scan.py +0 -0
  65. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/state_dict.py +0 -0
  66. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/_src/util.py +0 -0
  67. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/axis.py +0 -0
  68. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/core.py +0 -0
  69. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/debug.py +0 -0
  70. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/field.py +0 -0
  71. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/haxtyping.py +0 -0
  72. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/hof.py +0 -0
  73. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/jax_utils.py +0 -0
  74. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/__init__.py +0 -0
  75. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/activations.py +0 -0
  76. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/attention.py +0 -0
  77. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/conv.py +0 -0
  78. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/dropout.py +0 -0
  79. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/embedding.py +0 -0
  80. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/linear.py +0 -0
  81. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/loss.py +0 -0
  82. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/mlp.py +0 -0
  83. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/normalization.py +0 -0
  84. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/pool.py +0 -0
  85. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/nn/scan.py +0 -0
  86. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/partitioning.py +0 -0
  87. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/quantization.py +0 -0
  88. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/random.py +0 -0
  89. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/specialized_fns.py +0 -0
  90. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/state_dict.py +0 -0
  91. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/tree_util.py +0 -0
  92. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/types.py +0 -0
  93. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/util.py +0 -0
  94. {haliax-1.4.dev406 → haliax-1.4.dev408}/src/haliax/wrap.py +0 -0
  95. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/core_test.py +0 -0
  96. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_attention.py +0 -0
  97. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_axis.py +0 -0
  98. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_conv.py +0 -0
  99. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_debug.py +0 -0
  100. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_dot.py +0 -0
  101. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_dtype_typing.py +0 -0
  102. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_einsum.py +0 -0
  103. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_field.py +0 -0
  104. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_fp8.py +0 -0
  105. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_hof.py +0 -0
  106. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_int8.py +0 -0
  107. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_moe_linear.py +0 -0
  108. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_namedarray_typing.py +0 -0
  109. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_nn.py +0 -0
  110. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_parsing.py +0 -0
  111. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_partitioning.py +0 -0
  112. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_pool.py +0 -0
  113. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_random.py +0 -0
  114. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_rearrange.py +0 -0
  115. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_scan.py +0 -0
  116. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_scatter_gather.py +0 -0
  117. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_specialized_fns.py +0 -0
  118. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_state_dict.py +0 -0
  119. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_tree_util.py +0 -0
  120. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_utils.py +0 -0
  121. {haliax-1.4.dev406 → haliax-1.4.dev408}/tests/test_visualize_sharding.py +0 -0
  122. {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
- - [ ] `allclose`
8
- - [ ] `amin`
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
- - [ ] `array_equal`
15
- - [ ] `array_equiv`
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
- - [ ] `nanargmax`
97
- - [ ] `nanargmin`
98
- - [ ] `nancumprod`
99
- - [ ] `nancumsum`
100
- - [ ] `nanmax`
101
- - [ ] `nanmean`
96
+ - [x] `nanargmax`
97
+ - [x] `nanargmin`
98
+ - [x] `nancumprod`
99
+ - [x] `nancumsum`
100
+ - [x] `nanmax`
101
+ - [x] `nanmean`
102
102
  - [ ] `nanmedian`
103
- - [ ] `nanmin`
103
+ - [x] `nanmin`
104
104
  - [ ] `nanpercentile`
105
- - [ ] `nanprod`
105
+ - [x] `nanprod`
106
106
  - [ ] `nanquantile`
107
- - [ ] `nanstd`
108
- - [ ] `nansum`
109
- - [ ] `nanvar`
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.dev406
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