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