haliax 1.4.dev413__tar.gz → 1.4.dev420__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 (130) hide show
  1. {haliax-1.4.dev413 → haliax-1.4.dev420}/.agents/projects/api_parity.md +12 -12
  2. {haliax-1.4.dev413 → haliax-1.4.dev420}/.github/workflows/publish_dev.yaml +1 -1
  3. {haliax-1.4.dev413 → haliax-1.4.dev420}/.github/workflows/run_pre_commit.yaml +1 -1
  4. {haliax-1.4.dev413 → haliax-1.4.dev420}/.github/workflows/run_quick_levanter_tests.yaml +1 -1
  5. {haliax-1.4.dev413 → haliax-1.4.dev420}/.github/workflows/run_tests.yaml +2 -2
  6. {haliax-1.4.dev413 → haliax-1.4.dev420}/AUTHORS.md +1 -0
  7. {haliax-1.4.dev413 → haliax-1.4.dev420}/PKG-INFO +3 -2
  8. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/api.md +15 -0
  9. {haliax-1.4.dev413 → haliax-1.4.dev420}/pyproject.toml +3 -4
  10. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/__about__.py +1 -1
  11. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/__init__.py +113 -74
  12. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/dot.py +12 -13
  13. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/einsum.py +2 -3
  14. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/parsing.py +5 -5
  15. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/rearrange.py +7 -7
  16. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/scan.py +8 -7
  17. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/state_dict.py +16 -16
  18. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/axis.py +14 -14
  19. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/core.py +119 -102
  20. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/debug.py +5 -5
  21. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/fft.py +4 -1
  22. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/jax_utils.py +8 -8
  23. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/__init__.py +18 -7
  24. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/attention.py +14 -15
  25. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/conv.py +4 -4
  26. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/dropout.py +4 -6
  27. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/embedding.py +3 -4
  28. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/linear.py +7 -7
  29. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/loss.py +20 -21
  30. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/mlp.py +2 -2
  31. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/normalization.py +14 -14
  32. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/pool.py +5 -5
  33. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/scan.py +9 -11
  34. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/ops.py +8 -10
  35. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/partitioning.py +141 -100
  36. haliax-1.4.dev420/src/haliax/poly.py +304 -0
  37. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/quantization.py +4 -4
  38. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/random.py +2 -5
  39. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/specialized_fns.py +3 -5
  40. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/state_dict.py +2 -2
  41. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/tree_util.py +3 -2
  42. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/types.py +11 -11
  43. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/util.py +3 -3
  44. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/wrap.py +7 -7
  45. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/core_test.py +46 -0
  46. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_bitwise_ops.py +4 -0
  47. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_fft.py +5 -3
  48. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_moe_linear.py +3 -2
  49. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_nan_reductions.py +4 -0
  50. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_nn.py +18 -0
  51. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_partitioning.py +15 -27
  52. haliax-1.4.dev420/tests/test_poly_ops.py +134 -0
  53. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_utils.py +5 -2
  54. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_visualize_sharding.py +5 -2
  55. haliax-1.4.dev420/uv.lock +1555 -0
  56. haliax-1.4.dev413/uv.lock +0 -1711
  57. {haliax-1.4.dev413 → haliax-1.4.dev420}/.coveragerc +0 -0
  58. {haliax-1.4.dev413 → haliax-1.4.dev420}/.flake8 +0 -0
  59. {haliax-1.4.dev413 → haliax-1.4.dev420}/.gitignore +0 -0
  60. {haliax-1.4.dev413 → haliax-1.4.dev420}/.playbooks/add-types.md +0 -0
  61. {haliax-1.4.dev413 → haliax-1.4.dev420}/.playbooks/wrap-non-named.md +0 -0
  62. {haliax-1.4.dev413 → haliax-1.4.dev420}/.pre-commit-config.yaml +0 -0
  63. {haliax-1.4.dev413 → haliax-1.4.dev420}/.readthedocs.yaml +0 -0
  64. {haliax-1.4.dev413 → haliax-1.4.dev420}/AGENTS.md +0 -0
  65. {haliax-1.4.dev413 → haliax-1.4.dev420}/CONTRIBUTING.md +0 -0
  66. {haliax-1.4.dev413 → haliax-1.4.dev420}/CONTRIBUTORS.md +0 -0
  67. {haliax-1.4.dev413 → haliax-1.4.dev420}/LICENSE +0 -0
  68. {haliax-1.4.dev413 → haliax-1.4.dev420}/README.md +0 -0
  69. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/broadcasting.md +0 -0
  70. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/cheatsheet.md +0 -0
  71. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/css/material.css +0 -0
  72. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/css/mkdocstrings.css +0 -0
  73. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/faq.md +0 -0
  74. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/data_parallel_mesh.png +0 -0
  75. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  76. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_1d.png +0 -0
  77. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_1d_zero.png +0 -0
  78. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d.png +0 -0
  79. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  80. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  81. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  82. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  83. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/figures/device_mesh_2d_zero.png +0 -0
  84. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/fp8.md +0 -0
  85. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/index.md +0 -0
  86. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/indexing.md +0 -0
  87. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/matmul.md +0 -0
  88. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/nn.md +0 -0
  89. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/partitioning.md +0 -0
  90. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/primer.md +0 -0
  91. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/rearrange.ipynb +0 -0
  92. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/rearrange.md +0 -0
  93. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/requirements.txt +0 -0
  94. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/scan.md +0 -0
  95. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/state-dict.md +0 -0
  96. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/tutorial.md +0 -0
  97. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/typing.md +0 -0
  98. {haliax-1.4.dev413 → haliax-1.4.dev420}/docs/vmap.md +0 -0
  99. {haliax-1.4.dev413 → haliax-1.4.dev420}/etc/license_header.txt +0 -0
  100. {haliax-1.4.dev413 → haliax-1.4.dev420}/mkdocs.yml +0 -0
  101. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/__init__.py +0 -0
  102. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/compile_utils.py +0 -0
  103. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/fp8.py +0 -0
  104. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/_src/util.py +0 -0
  105. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/field.py +0 -0
  106. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/haxtyping.py +0 -0
  107. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/hof.py +0 -0
  108. {haliax-1.4.dev413 → haliax-1.4.dev420}/src/haliax/nn/activations.py +0 -0
  109. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_attention.py +0 -0
  110. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_axis.py +0 -0
  111. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_conv.py +0 -0
  112. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_debug.py +0 -0
  113. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_dot.py +0 -0
  114. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_dtype_typing.py +0 -0
  115. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_einsum.py +0 -0
  116. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_field.py +0 -0
  117. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_fp8.py +0 -0
  118. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_hof.py +0 -0
  119. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_int8.py +0 -0
  120. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_namedarray_typing.py +0 -0
  121. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_ops.py +0 -0
  122. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_parsing.py +0 -0
  123. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_pool.py +0 -0
  124. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_random.py +0 -0
  125. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_rearrange.py +0 -0
  126. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_scan.py +0 -0
  127. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_scatter_gather.py +0 -0
  128. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_specialized_fns.py +0 -0
  129. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_state_dict.py +0 -0
  130. {haliax-1.4.dev413 → haliax-1.4.dev420}/tests/test_tree_util.py +0 -0
@@ -115,15 +115,15 @@ APIs that don't translate well to named tensors are intentionally omitted here.
115
115
  - [ ] `permute_dims`
116
116
  - [ ] `piecewise`
117
117
  - [ ] `place`
118
- - [ ] `poly`
119
- - [ ] `polyadd`
120
- - [ ] `polyder`
121
- - [ ] `polydiv`
122
- - [ ] `polyfit`
123
- - [ ] `polyint`
124
- - [ ] `polymul`
125
- - [ ] `polysub`
126
- - [ ] `polyval`
118
+ - [x] `poly`
119
+ - [x] `polyadd`
120
+ - [x] `polyder`
121
+ - [x] `polydiv`
122
+ - [x] `polyfit`
123
+ - [x] `polyint`
124
+ - [x] `polymul`
125
+ - [x] `polysub`
126
+ - [x] `polyval`
127
127
  - [ ] `pow`
128
128
  - [ ] `promote_types`
129
129
  - [ ] `put`
@@ -134,7 +134,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
134
134
  - [ ] `resize`
135
135
  - [ ] `result_type`
136
136
  - [ ] `rollaxis`
137
- - [ ] `roots`
137
+ - [x] `roots`
138
138
  - [ ] `rot90`
139
139
  - [ ] `select`
140
140
  - [ ] `setdiff1d`
@@ -148,7 +148,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
148
148
  - [ ] `tri`
149
149
  - [ ] `tril_indices`
150
150
  - [ ] `tril_indices_from`
151
- - [ ] `trim_zeros`
151
+ - [x] `trim_zeros`
152
152
  - [ ] `triu_indices`
153
153
  - [ ] `triu_indices_from`
154
154
  - [ ] `union1d`
@@ -156,7 +156,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
156
156
  - [ ] `unravel_index`
157
157
  - [ ] `unstack`
158
158
  - [ ] `unwrap`
159
- - [ ] `vander`
159
+ - [x] `vander`
160
160
  - [ ] `vsplit`
161
161
  - [ ] `vstack`
162
162
 
@@ -20,7 +20,7 @@ jobs:
20
20
  - name: Set up Python
21
21
  uses: actions/setup-python@v2
22
22
  with:
23
- python-version: '3.x'
23
+ python-version: '3.11'
24
24
 
25
25
  - name: Calculate Version and Build Number
26
26
  run: |
@@ -8,7 +8,7 @@ jobs:
8
8
  runs-on: ubuntu-latest
9
9
  strategy:
10
10
  matrix:
11
- python-version: ["3.10.11"]
11
+ python-version: ["3.11"]
12
12
 
13
13
  steps:
14
14
  - uses: actions/checkout@v3
@@ -9,7 +9,7 @@ jobs:
9
9
  runs-on: ubuntu-latest
10
10
  strategy:
11
11
  matrix:
12
- python-version: ["3.10", "3.11"]
12
+ python-version: ["3.11"]
13
13
  steps:
14
14
  - name: Checkout repository
15
15
  uses: actions/checkout@v3
@@ -9,10 +9,10 @@ jobs:
9
9
 
10
10
  steps:
11
11
  - uses: actions/checkout@v3
12
- - name: Set up Python 3.10.11
12
+ - name: Set up Python 3.11
13
13
  uses: actions/setup-python@v4
14
14
  with:
15
- python-version: 3.10.11
15
+ python-version: 3.11
16
16
  - name: Install dependencies
17
17
  run: |
18
18
  python -m pip install uv
@@ -3,3 +3,4 @@
3
3
  The Levanter Authors currently include:
4
4
 
5
5
  - The Board of Trustees of the Leland Stanford Junior University
6
+ - Open Athena AI Foundation, Inc.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev413
3
+ Version: 1.4.dev420
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/
@@ -14,9 +14,10 @@ Classifier: License :: OSI Approved :: Apache Software License
14
14
  Classifier: Operating System :: MacOS :: MacOS X
15
15
  Classifier: Operating System :: POSIX :: Linux
16
16
  Classifier: Programming Language :: Python :: 3
17
- Requires-Python: >=3.10
17
+ Requires-Python: >=3.11
18
18
  Requires-Dist: aqtp>=0.8.2
19
19
  Requires-Dist: equinox>=0.10.6
20
+ Requires-Dist: jax>=0.6.2
20
21
  Requires-Dist: jaxtyping>=0.2.20
21
22
  Requires-Dist: jmp>=0.0.4
22
23
  Requires-Dist: safetensors>=0.4.3
@@ -340,6 +340,21 @@ These are all more or less directly from JAX's NumPy API.
340
340
  ::: haliax.subtract
341
341
  ::: haliax.true_divide
342
342
 
343
+ ### Polynomial Operations
344
+
345
+ ::: haliax.poly
346
+ ::: haliax.polyadd
347
+ ::: haliax.polysub
348
+ ::: haliax.polymul
349
+ ::: haliax.polydiv
350
+ ::: haliax.polyint
351
+ ::: haliax.polyder
352
+ ::: haliax.polyval
353
+ ::: haliax.polyfit
354
+ ::: haliax.roots
355
+ ::: haliax.trim_zeros
356
+ ::: haliax.vander
357
+
343
358
  ### Other Operations
344
359
 
345
360
  ::: haliax.bincount
@@ -11,7 +11,7 @@ authors = [
11
11
  ]
12
12
  description = "Named Tensors for Legible Deep Learning in JAX"
13
13
  readme = "README.md"
14
- requires-python = ">=3.10"
14
+ requires-python = ">=3.11"
15
15
  classifiers = [
16
16
  "Programming Language :: Python :: 3",
17
17
  "License :: OSI Approved :: Apache Software License",
@@ -21,8 +21,7 @@ classifiers = [
21
21
  "Intended Audience :: Science/Research",
22
22
  ]
23
23
  dependencies = [
24
- # we require that you install jax yourself, since the extras vary by system.
25
- # jax = {version = ">=0.4.19,<0.5.0"}
24
+ "jax >= 0.6.2",
26
25
  "equinox>=0.10.6",
27
26
  "jaxtyping>=0.2.20",
28
27
  "jmp>=0.0.4",
@@ -61,7 +60,7 @@ haliax = ["src/haliax/*"]
61
60
 
62
61
  [tool.black]
63
62
  line-length = 119
64
- target-version = ["py310"]
63
+ target-version = ["py311"]
65
64
  preview = true
66
65
 
67
66
  [tool.isort]
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
 
6
- __version__ = "1.4.dev413"
6
+ __version__ = "1.4.dev420"
@@ -4,15 +4,12 @@
4
4
 
5
5
 
6
6
  import typing as t
7
- from typing import Optional, Sequence
7
+ from typing import Sequence
8
8
 
9
9
  import jax
10
10
  import jax.numpy as jnp
11
+ from jax.typing import DTypeLike
11
12
 
12
- try:
13
- from jax.typing import DTypeLike
14
- except ImportError:
15
- from jax._src.typing import DTypeLike
16
13
 
17
14
  import haliax.debug as debug
18
15
  import haliax.nn as nn
@@ -96,6 +93,22 @@ from .ops import (
96
93
  bincount,
97
94
  where,
98
95
  )
96
+
97
+ from .poly import (
98
+ poly,
99
+ polyadd,
100
+ polysub,
101
+ polymul,
102
+ polydiv,
103
+ polyint,
104
+ polyder,
105
+ polyval,
106
+ polyfit,
107
+ roots,
108
+ trim_zeros,
109
+ vander,
110
+ )
111
+
99
112
  from .fft import (
100
113
  fft,
101
114
  fftfreq,
@@ -108,7 +121,7 @@ from .fft import (
108
121
  rfft,
109
122
  rfftfreq,
110
123
  )
111
- from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
124
+ from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, set_mesh, shard, shard_with_axis_mapping
112
125
  from .specialized_fns import top_k
113
126
  from .types import Scalar
114
127
  from .util import is_named_array
@@ -126,21 +139,21 @@ A = t.TypeVar("A", Scalar, NamedArray, jnp.ndarray)
126
139
 
127
140
 
128
141
  # creation routines
129
- def zeros(shape: AxisSpec, dtype: Optional[DTypeLike] = None) -> NamedArray:
142
+ def zeros(shape: AxisSpec, dtype: DTypeLike | None = None) -> NamedArray:
130
143
  """Creates a NamedArray with all elements set to 0"""
131
144
  if dtype is None:
132
145
  dtype = jnp.float32
133
146
  return full(shape, 0, dtype)
134
147
 
135
148
 
136
- def ones(shape: AxisSpec, dtype: Optional[DTypeLike] = None) -> NamedArray:
149
+ def ones(shape: AxisSpec, dtype: DTypeLike | None = None) -> NamedArray:
137
150
  """Creates a NamedArray with all elements set to 1"""
138
151
  if dtype is None:
139
152
  dtype = jnp.float32
140
153
  return full(shape, 1, dtype)
141
154
 
142
155
 
143
- def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> NamedArray:
156
+ def full(shape: AxisSpec, fill_value: T, dtype: DTypeLike | None = None) -> NamedArray:
144
157
  """Creates a NamedArray with all elements set to `fill_value`"""
145
158
  if isinstance(shape, Axis):
146
159
  return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
@@ -159,12 +172,12 @@ def ones_like(a: NamedArray, dtype=None) -> NamedArray:
159
172
  return NamedArray(jnp.ones_like(a.array, dtype=dtype), a.axes)
160
173
 
161
174
 
162
- def full_like(a: NamedArray, fill_value: T, dtype: Optional[DTypeLike] = None) -> NamedArray:
175
+ def full_like(a: NamedArray, fill_value: T, dtype: DTypeLike | None = None) -> NamedArray:
163
176
  """Creates a NamedArray with all elements set to `fill_value`"""
164
177
  return NamedArray(jnp.full_like(a.array, fill_value, dtype=dtype), a.axes)
165
178
 
166
179
 
167
- def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None) -> NamedArray:
180
+ def arange(axis: AxisSpec, *, start=0, step=1, dtype: DTypeLike | None = None) -> NamedArray:
168
181
  """
169
182
  Version of jnp.arange that returns a NamedArray.
170
183
 
@@ -195,7 +208,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
195
208
 
196
209
  # TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
197
210
  def linspace(
198
- axis: AxisSelector, *, start: float, stop: float, endpoint: bool = True, dtype: Optional[DTypeLike] = None
211
+ axis: AxisSelector, *, start: float, stop: float, endpoint: bool = True, dtype: DTypeLike | None = None
199
212
  ) -> NamedArray:
200
213
  """
201
214
  Version of jnp.linspace that returns a NamedArray.
@@ -213,7 +226,7 @@ def logspace(
213
226
  stop: float,
214
227
  endpoint: bool = True,
215
228
  base: float = 10.0,
216
- dtype: Optional[DTypeLike] = None,
229
+ dtype: DTypeLike | None = None,
217
230
  ) -> NamedArray:
218
231
  """
219
232
  Version of jnp.logspace that returns a NamedArray.
@@ -225,7 +238,7 @@ def logspace(
225
238
 
226
239
 
227
240
  def geomspace(
228
- axis: AxisSelector, *, start: float, stop: float, endpoint: bool = True, dtype: Optional[DTypeLike] = None
241
+ axis: AxisSelector, *, start: float, stop: float, endpoint: bool = True, dtype: DTypeLike | None = None
229
242
  ) -> NamedArray:
230
243
  """
231
244
  Version of jnp.geomspace that returns a NamedArray.
@@ -247,7 +260,7 @@ def stack(axis: AxisSelector, arrays: Sequence[NamedArray]) -> NamedArray:
247
260
 
248
261
 
249
262
  def repeat(
250
- a: NamedArray, repeats: int | jnp.ndarray, axis: AxisSelector, total_repeat_length: Optional[int] = None
263
+ a: NamedArray, repeats: int | jnp.ndarray, axis: AxisSelector, total_repeat_length: int | None = None
251
264
  ) -> NamedArray:
252
265
  """Version of [jax.numpy.repeat][] that returns a NamedArray"""
253
266
  index = a.axis_indices(axis)
@@ -574,91 +587,91 @@ def trunc(a: A) -> A:
574
587
 
575
588
 
576
589
  # Reduction functions
577
- def all(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
590
+ def all(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
578
591
  """
579
592
  Named version of [jax.numpy.all](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.all.html#jax.numpy.all).
580
593
  """
581
594
  return wrap_reduction_call(jnp.all, array, axis, where, single_axis_only=False, supports_where=True)
582
595
 
583
596
 
584
- def amax(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
597
+ def amax(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
585
598
  """
586
599
  Aliax for max. See max for details.
587
600
  """
588
601
  return wrap_reduction_call(jnp.amax, array, axis, where, single_axis_only=False, supports_where=True)
589
602
 
590
603
 
591
- def amin(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
604
+ def amin(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
592
605
  """
593
606
  Aliax for min. See min for details.
594
607
  """
595
608
  return wrap_reduction_call(jnp.amin, array, axis, where, single_axis_only=False, supports_where=True)
596
609
 
597
610
 
598
- def any(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
611
+ def any(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
599
612
  """True if any elements along a given axis or axes are True. If axis is None, any elements are True."""
600
613
  return wrap_reduction_call(jnp.any, array, axis, where, single_axis_only=False, supports_where=True)
601
614
 
602
615
 
603
- def argmax(array: NamedArray, axis: Optional[AxisSelector]) -> NamedArray:
616
+ def argmax(array: NamedArray, axis: AxisSelector | None) -> NamedArray:
604
617
  return wrap_reduction_call(jnp.argmax, array, axis, None, single_axis_only=True, supports_where=False)
605
618
 
606
619
 
607
- def argmin(array: NamedArray, axis: Optional[AxisSelector]) -> NamedArray:
620
+ def argmin(array: NamedArray, axis: AxisSelector | None) -> NamedArray:
608
621
  return wrap_reduction_call(jnp.argmin, array, axis, None, single_axis_only=True, supports_where=False)
609
622
 
610
623
 
611
- def max(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
624
+ def max(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
612
625
  return wrap_reduction_call(jnp.max, array, axis, where, single_axis_only=False, supports_where=True)
613
626
 
614
627
 
615
628
  def mean(
616
629
  array: NamedArray,
617
- axis: Optional[AxisSelection] = None,
630
+ axis: AxisSelection | None = None,
618
631
  *,
619
- where: Optional[NamedArray] = None,
620
- dtype: Optional[DTypeLike] = None,
632
+ where: NamedArray | None = None,
633
+ dtype: DTypeLike | None = None,
621
634
  ) -> NamedArray:
622
635
  return wrap_reduction_call(jnp.mean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
623
636
 
624
637
 
625
- def min(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
638
+ def min(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
626
639
  return wrap_reduction_call(jnp.min, array, axis, where, single_axis_only=False, supports_where=True)
627
640
 
628
641
 
629
642
  def prod(
630
643
  array: NamedArray,
631
- axis: Optional[AxisSelection] = None,
644
+ axis: AxisSelection | None = None,
632
645
  *,
633
- where: Optional[NamedArray] = None,
634
- dtype: Optional[DTypeLike] = None,
646
+ where: NamedArray | None = None,
647
+ dtype: DTypeLike | None = None,
635
648
  ) -> NamedArray:
636
649
  return wrap_reduction_call(jnp.prod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
637
650
 
638
651
 
639
652
  def std(
640
653
  array: NamedArray,
641
- axis: Optional[AxisSelection] = None,
654
+ axis: AxisSelection | None = None,
642
655
  *,
643
- where: Optional[NamedArray] = None,
656
+ where: NamedArray | None = None,
644
657
  ddof: int = 0,
645
- dtype: Optional[DTypeLike] = None,
658
+ dtype: DTypeLike | None = None,
646
659
  ) -> NamedArray:
647
660
  return wrap_reduction_call(
648
661
  jnp.std, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
649
662
  )
650
663
 
651
664
 
652
- def ptp(array: NamedArray, axis: Optional[AxisSelection] = None, *, where: Optional[NamedArray] = None) -> NamedArray:
665
+ def ptp(array: NamedArray, axis: AxisSelection | None = None, *, where: NamedArray | None = None) -> NamedArray:
653
666
  return wrap_reduction_call(jnp.ptp, array, axis, where, single_axis_only=False, supports_where=True)
654
667
 
655
668
 
656
669
  def product(
657
670
  array: NamedArray,
658
- axis: Optional[AxisSelection] = None,
671
+ axis: AxisSelection | None = None,
659
672
  *,
660
- where: Optional[NamedArray] = None,
661
- dtype: Optional[DTypeLike] = None,
673
+ where: NamedArray | None = None,
674
+ dtype: DTypeLike | None = None,
662
675
  ) -> NamedArray:
663
676
  return wrap_reduction_call(
664
677
  jnp.product, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
@@ -670,130 +683,140 @@ _sum = sum
670
683
 
671
684
  def sum(
672
685
  array: NamedArray,
673
- axis: Optional[AxisSelection] = None,
686
+ axis: AxisSelection | None = None,
674
687
  *,
675
- where: Optional[NamedArray] = None,
676
- dtype: Optional[DTypeLike] = None,
688
+ where: NamedArray | None = None,
689
+ dtype: DTypeLike | None = None,
677
690
  ) -> NamedArray:
678
691
  return wrap_reduction_call(jnp.sum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
679
692
 
680
693
 
681
694
  def var(
682
695
  array: NamedArray,
683
- axis: Optional[AxisSelection] = None,
696
+ axis: AxisSelection | None = None,
684
697
  *,
685
- where: Optional[NamedArray] = None,
698
+ where: NamedArray | None = None,
686
699
  ddof: int = 0,
687
- dtype: Optional[DTypeLike] = None,
700
+ dtype: DTypeLike | None = None,
688
701
  ) -> NamedArray:
689
702
  return wrap_reduction_call(
690
703
  jnp.var, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
691
704
  )
692
705
 
693
706
 
694
- def nanargmax(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
707
+ def nanargmax(array: NamedArray, axis: AxisSelector | None = None) -> NamedArray:
695
708
  return wrap_reduction_call(jnp.nanargmax, array, axis, None, single_axis_only=True, supports_where=False)
696
709
 
697
710
 
698
- def nanargmin(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
711
+ def nanargmin(array: NamedArray, axis: AxisSelector | None = None) -> NamedArray:
699
712
  return wrap_reduction_call(jnp.nanargmin, array, axis, None, single_axis_only=True, supports_where=False)
700
713
 
701
714
 
702
715
  def nanmax(
703
716
  array: NamedArray,
704
- axis: Optional[AxisSelection] = None,
717
+ axis: AxisSelection | None = None,
705
718
  *,
706
- where: Optional[NamedArray] = None,
719
+ where: NamedArray | None = None,
707
720
  ) -> NamedArray:
708
721
  return wrap_reduction_call(jnp.nanmax, array, axis, where, single_axis_only=False, supports_where=True)
709
722
 
710
723
 
711
724
  def nanmean(
712
725
  array: NamedArray,
713
- axis: Optional[AxisSelection] = None,
726
+ axis: AxisSelection | None = None,
714
727
  *,
715
- where: Optional[NamedArray] = None,
716
- dtype: Optional[DTypeLike] = None,
728
+ where: NamedArray | None = None,
729
+ dtype: DTypeLike | None = None,
717
730
  ) -> NamedArray:
718
- return wrap_reduction_call(jnp.nanmean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
731
+ return wrap_reduction_call(
732
+ jnp.nanmean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
733
+ )
719
734
 
720
735
 
721
736
  def nanmin(
722
737
  array: NamedArray,
723
- axis: Optional[AxisSelection] = None,
738
+ axis: AxisSelection | None = None,
724
739
  *,
725
- where: Optional[NamedArray] = None,
740
+ where: NamedArray | None = None,
726
741
  ) -> NamedArray:
727
742
  return wrap_reduction_call(jnp.nanmin, array, axis, where, single_axis_only=False, supports_where=True)
728
743
 
729
744
 
730
745
  def nanprod(
731
746
  array: NamedArray,
732
- axis: Optional[AxisSelection] = None,
747
+ axis: AxisSelection | None = None,
733
748
  *,
734
- where: Optional[NamedArray] = None,
735
- dtype: Optional[DTypeLike] = None,
749
+ where: NamedArray | None = None,
750
+ dtype: DTypeLike | None = None,
736
751
  ) -> NamedArray:
737
- return wrap_reduction_call(jnp.nanprod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
752
+ return wrap_reduction_call(
753
+ jnp.nanprod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
754
+ )
738
755
 
739
756
 
740
757
  def nanstd(
741
758
  array: NamedArray,
742
- axis: Optional[AxisSelection] = None,
759
+ axis: AxisSelection | None = None,
743
760
  *,
744
- where: Optional[NamedArray] = None,
761
+ where: NamedArray | None = None,
745
762
  ddof: int = 0,
746
- dtype: Optional[DTypeLike] = None,
763
+ dtype: DTypeLike | None = None,
747
764
  ) -> NamedArray:
748
- return wrap_reduction_call(jnp.nanstd, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof)
765
+ return wrap_reduction_call(
766
+ jnp.nanstd, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
767
+ )
749
768
 
750
769
 
751
770
  def nansum(
752
771
  array: NamedArray,
753
- axis: Optional[AxisSelection] = None,
772
+ axis: AxisSelection | None = None,
754
773
  *,
755
- where: Optional[NamedArray] = None,
756
- dtype: Optional[DTypeLike] = None,
774
+ where: NamedArray | None = None,
775
+ dtype: DTypeLike | None = None,
757
776
  ) -> NamedArray:
758
- return wrap_reduction_call(jnp.nansum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
777
+ return wrap_reduction_call(
778
+ jnp.nansum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
779
+ )
759
780
 
760
781
 
761
782
  def nanvar(
762
783
  array: NamedArray,
763
- axis: Optional[AxisSelection] = None,
784
+ axis: AxisSelection | None = None,
764
785
  *,
765
- where: Optional[NamedArray] = None,
786
+ where: NamedArray | None = None,
766
787
  ddof: int = 0,
767
- dtype: Optional[DTypeLike] = None,
788
+ dtype: DTypeLike | None = None,
768
789
  ) -> NamedArray:
769
- return wrap_reduction_call(jnp.nanvar, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof)
790
+ return wrap_reduction_call(
791
+ jnp.nanvar, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
792
+ )
770
793
 
771
794
 
772
795
  # "Normalization" functions that use an axis but don't change the shape
773
796
 
774
797
 
775
- def cumsum(a: NamedArray, axis: AxisSelector, *, dtype: Optional[DTypeLike] = None) -> NamedArray:
798
+ def cumsum(a: NamedArray, axis: AxisSelector, *, dtype: DTypeLike | None = None) -> NamedArray:
776
799
  """
777
800
  Named version of [jax.numpy.cumsum](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.cumsum.html)
778
801
  """
779
802
  return wrap_axiswise_call(jnp.cumsum, a, axis, dtype=dtype, single_axis_only=True)
780
803
 
781
804
 
782
- def cumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None) -> NamedArray:
805
+ def cumprod(a: NamedArray, axis: AxisSelector, dtype: DTypeLike | None = None) -> NamedArray:
783
806
  """
784
807
  Named version of [jax.numpy.cumprod](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.cumprod.html)
785
808
  """
786
809
  return wrap_axiswise_call(jnp.cumprod, a, axis, dtype=dtype, single_axis_only=True)
787
810
 
788
811
 
789
- def nancumsum(a: NamedArray, axis: AxisSelector, *, dtype: Optional[DTypeLike] = None) -> NamedArray:
812
+ def nancumsum(a: NamedArray, axis: AxisSelector, *, dtype: DTypeLike | None = None) -> NamedArray:
790
813
  """
791
814
  Named version of [jax.numpy.nancumsum](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumsum.html)
792
815
  """
793
816
  return wrap_axiswise_call(jnp.nancumsum, a, axis, dtype=dtype, single_axis_only=True)
794
817
 
795
818
 
796
- def nancumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None) -> NamedArray:
819
+ def nancumprod(a: NamedArray, axis: AxisSelector, dtype: DTypeLike | None = None) -> NamedArray:
797
820
  """
798
821
  Named version of [jax.numpy.nancumprod](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumprod.html)
799
822
  """
@@ -807,14 +830,17 @@ def sort(a: NamedArray, axis: AxisSelector) -> NamedArray:
807
830
  return wrap_axiswise_call(jnp.sort, a, axis, single_axis_only=True)
808
831
 
809
832
 
810
- def argsort(a: NamedArray, axis: AxisSelector) -> NamedArray:
833
+ def argsort(a: NamedArray, axis: AxisSelector | None, *, stable: bool = False) -> NamedArray:
811
834
  """
812
835
  Named version of [jax.numpy.argsort](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.argsort.html).
813
836
 
814
837
  If `axis` is None, the returned array will be a 1D array of indices that would sort the flattened array,
815
838
  identical to `jax.numpy.argsort(a.array)`.
839
+
840
+ Args:
841
+ stable: If ``True``, ensures that the indices of equal elements preserve their relative order.
816
842
  """
817
- return wrap_axiswise_call(jnp.argsort, a, axis, single_axis_only=True)
843
+ return wrap_axiswise_call(jnp.argsort, a, axis, single_axis_only=True, stable=stable)
818
844
 
819
845
 
820
846
  # elemwise binary ops
@@ -1226,6 +1252,18 @@ __all__ = [
1226
1252
  "clip",
1227
1253
  "tril",
1228
1254
  "triu",
1255
+ "poly",
1256
+ "polyadd",
1257
+ "polysub",
1258
+ "polymul",
1259
+ "polydiv",
1260
+ "polyint",
1261
+ "polyder",
1262
+ "polyval",
1263
+ "polyfit",
1264
+ "roots",
1265
+ "trim_zeros",
1266
+ "vander",
1229
1267
  "fft",
1230
1268
  "ifft",
1231
1269
  "hfft",
@@ -1311,4 +1349,5 @@ __all__ = [
1311
1349
  "NamedArrayAxes",
1312
1350
  "NamedArrayAxesSpec",
1313
1351
  "Named",
1352
+ "set_mesh",
1314
1353
  ]