haliax 1.4.dev419__tar.gz → 1.4.dev438__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 (133) hide show
  1. {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/publish_dev.yaml +1 -1
  2. {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_pre_commit.yaml +1 -1
  3. {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_quick_levanter_tests.yaml +1 -1
  4. {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_tests.yaml +2 -2
  5. {haliax-1.4.dev419 → haliax-1.4.dev438}/PKG-INFO +2 -2
  6. {haliax-1.4.dev419 → haliax-1.4.dev438}/pyproject.toml +2 -2
  7. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/__about__.py +1 -1
  8. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/__init__.py +62 -62
  9. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/dot.py +12 -13
  10. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/einsum.py +2 -3
  11. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/parsing.py +5 -5
  12. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/rearrange.py +7 -7
  13. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/scan.py +8 -7
  14. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/state_dict.py +15 -15
  15. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/axis.py +14 -14
  16. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/core.py +95 -98
  17. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/debug.py +5 -5
  18. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/jax_utils.py +8 -8
  19. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/attention.py +14 -15
  20. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/conv.py +4 -4
  21. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/dropout.py +4 -6
  22. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/embedding.py +24 -11
  23. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/linear.py +56 -12
  24. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/loss.py +20 -21
  25. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/mlp.py +2 -2
  26. haliax-1.4.dev438/src/haliax/nn/mup.py +206 -0
  27. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/normalization.py +14 -14
  28. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/pool.py +5 -5
  29. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/scan.py +9 -11
  30. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/ops.py +6 -6
  31. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/partitioning.py +44 -42
  32. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/quantization.py +4 -4
  33. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/random.py +2 -5
  34. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/specialized_fns.py +3 -5
  35. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/state_dict.py +2 -2
  36. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/tree_util.py +3 -2
  37. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/types.py +11 -11
  38. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/util.py +3 -3
  39. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/wrap.py +7 -7
  40. haliax-1.4.dev438/tests/test_mup_coordinate_check.py +164 -0
  41. haliax-1.4.dev438/tests/test_mup_embedding.py +48 -0
  42. haliax-1.4.dev438/tests/test_mup_linear.py +120 -0
  43. {haliax-1.4.dev419 → haliax-1.4.dev438}/uv.lock +43 -439
  44. {haliax-1.4.dev419 → haliax-1.4.dev438}/.agents/projects/api_parity.md +0 -0
  45. {haliax-1.4.dev419 → haliax-1.4.dev438}/.coveragerc +0 -0
  46. {haliax-1.4.dev419 → haliax-1.4.dev438}/.flake8 +0 -0
  47. {haliax-1.4.dev419 → haliax-1.4.dev438}/.gitignore +0 -0
  48. {haliax-1.4.dev419 → haliax-1.4.dev438}/.playbooks/add-types.md +0 -0
  49. {haliax-1.4.dev419 → haliax-1.4.dev438}/.playbooks/wrap-non-named.md +0 -0
  50. {haliax-1.4.dev419 → haliax-1.4.dev438}/.pre-commit-config.yaml +0 -0
  51. {haliax-1.4.dev419 → haliax-1.4.dev438}/.readthedocs.yaml +0 -0
  52. {haliax-1.4.dev419 → haliax-1.4.dev438}/AGENTS.md +0 -0
  53. {haliax-1.4.dev419 → haliax-1.4.dev438}/AUTHORS.md +0 -0
  54. {haliax-1.4.dev419 → haliax-1.4.dev438}/CONTRIBUTING.md +0 -0
  55. {haliax-1.4.dev419 → haliax-1.4.dev438}/CONTRIBUTORS.md +0 -0
  56. {haliax-1.4.dev419 → haliax-1.4.dev438}/LICENSE +0 -0
  57. {haliax-1.4.dev419 → haliax-1.4.dev438}/README.md +0 -0
  58. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/api.md +0 -0
  59. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/broadcasting.md +0 -0
  60. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/cheatsheet.md +0 -0
  61. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/css/material.css +0 -0
  62. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/css/mkdocstrings.css +0 -0
  63. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/faq.md +0 -0
  64. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh.png +0 -0
  65. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  66. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_1d.png +0 -0
  67. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_1d_zero.png +0 -0
  68. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d.png +0 -0
  69. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  70. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  71. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  72. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  73. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_zero.png +0 -0
  74. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/fp8.md +0 -0
  75. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/index.md +0 -0
  76. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/indexing.md +0 -0
  77. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/matmul.md +0 -0
  78. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/nn.md +0 -0
  79. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/partitioning.md +0 -0
  80. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/primer.md +0 -0
  81. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/rearrange.ipynb +0 -0
  82. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/rearrange.md +0 -0
  83. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/requirements.txt +0 -0
  84. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/scan.md +0 -0
  85. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/state-dict.md +0 -0
  86. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/tutorial.md +0 -0
  87. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/typing.md +0 -0
  88. {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/vmap.md +0 -0
  89. {haliax-1.4.dev419 → haliax-1.4.dev438}/etc/license_header.txt +0 -0
  90. {haliax-1.4.dev419 → haliax-1.4.dev438}/mkdocs.yml +0 -0
  91. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/__init__.py +0 -0
  92. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/compile_utils.py +0 -0
  93. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/fp8.py +0 -0
  94. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/util.py +0 -0
  95. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/fft.py +0 -0
  96. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/field.py +0 -0
  97. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/haxtyping.py +0 -0
  98. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/hof.py +0 -0
  99. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/__init__.py +0 -0
  100. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/activations.py +0 -0
  101. {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/poly.py +0 -0
  102. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/core_test.py +0 -0
  103. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_attention.py +0 -0
  104. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_axis.py +0 -0
  105. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_bitwise_ops.py +0 -0
  106. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_conv.py +0 -0
  107. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_debug.py +0 -0
  108. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_dot.py +0 -0
  109. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_dtype_typing.py +0 -0
  110. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_einsum.py +0 -0
  111. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_fft.py +0 -0
  112. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_field.py +0 -0
  113. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_fp8.py +0 -0
  114. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_hof.py +0 -0
  115. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_int8.py +0 -0
  116. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_moe_linear.py +0 -0
  117. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_namedarray_typing.py +0 -0
  118. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_nan_reductions.py +0 -0
  119. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_nn.py +0 -0
  120. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_ops.py +0 -0
  121. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_parsing.py +0 -0
  122. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_partitioning.py +0 -0
  123. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_poly_ops.py +0 -0
  124. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_pool.py +0 -0
  125. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_random.py +0 -0
  126. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_rearrange.py +0 -0
  127. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_scan.py +0 -0
  128. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_scatter_gather.py +0 -0
  129. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_specialized_fns.py +0 -0
  130. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_state_dict.py +0 -0
  131. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_tree_util.py +0 -0
  132. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_utils.py +0 -0
  133. {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_visualize_sharding.py +0 -0
@@ -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
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev419
3
+ Version: 1.4.dev438
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,7 +14,7 @@ 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
20
  Requires-Dist: jax>=0.6.2
@@ -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",
@@ -60,7 +60,7 @@ haliax = ["src/haliax/*"]
60
60
 
61
61
  [tool.black]
62
62
  line-length = 119
63
- target-version = ["py310"]
63
+ target-version = ["py311"]
64
64
  preview = true
65
65
 
66
66
  [tool.isort]
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
 
6
- __version__ = "1.4.dev419"
6
+ __version__ = "1.4.dev438"
@@ -4,7 +4,7 @@
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
@@ -139,21 +139,21 @@ A = t.TypeVar("A", Scalar, NamedArray, jnp.ndarray)
139
139
 
140
140
 
141
141
  # creation routines
142
- def zeros(shape: AxisSpec, dtype: Optional[DTypeLike] = None) -> NamedArray:
142
+ def zeros(shape: AxisSpec, dtype: DTypeLike | None = None) -> NamedArray:
143
143
  """Creates a NamedArray with all elements set to 0"""
144
144
  if dtype is None:
145
145
  dtype = jnp.float32
146
146
  return full(shape, 0, dtype)
147
147
 
148
148
 
149
- def ones(shape: AxisSpec, dtype: Optional[DTypeLike] = None) -> NamedArray:
149
+ def ones(shape: AxisSpec, dtype: DTypeLike | None = None) -> NamedArray:
150
150
  """Creates a NamedArray with all elements set to 1"""
151
151
  if dtype is None:
152
152
  dtype = jnp.float32
153
153
  return full(shape, 1, dtype)
154
154
 
155
155
 
156
- 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:
157
157
  """Creates a NamedArray with all elements set to `fill_value`"""
158
158
  if isinstance(shape, Axis):
159
159
  return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
@@ -172,12 +172,12 @@ def ones_like(a: NamedArray, dtype=None) -> NamedArray:
172
172
  return NamedArray(jnp.ones_like(a.array, dtype=dtype), a.axes)
173
173
 
174
174
 
175
- 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:
176
176
  """Creates a NamedArray with all elements set to `fill_value`"""
177
177
  return NamedArray(jnp.full_like(a.array, fill_value, dtype=dtype), a.axes)
178
178
 
179
179
 
180
- 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:
181
181
  """
182
182
  Version of jnp.arange that returns a NamedArray.
183
183
 
@@ -208,7 +208,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
208
208
 
209
209
  # TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
210
210
  def linspace(
211
- 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
212
212
  ) -> NamedArray:
213
213
  """
214
214
  Version of jnp.linspace that returns a NamedArray.
@@ -226,7 +226,7 @@ def logspace(
226
226
  stop: float,
227
227
  endpoint: bool = True,
228
228
  base: float = 10.0,
229
- dtype: Optional[DTypeLike] = None,
229
+ dtype: DTypeLike | None = None,
230
230
  ) -> NamedArray:
231
231
  """
232
232
  Version of jnp.logspace that returns a NamedArray.
@@ -238,7 +238,7 @@ def logspace(
238
238
 
239
239
 
240
240
  def geomspace(
241
- 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
242
242
  ) -> NamedArray:
243
243
  """
244
244
  Version of jnp.geomspace that returns a NamedArray.
@@ -260,7 +260,7 @@ def stack(axis: AxisSelector, arrays: Sequence[NamedArray]) -> NamedArray:
260
260
 
261
261
 
262
262
  def repeat(
263
- 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
264
264
  ) -> NamedArray:
265
265
  """Version of [jax.numpy.repeat][] that returns a NamedArray"""
266
266
  index = a.axis_indices(axis)
@@ -587,91 +587,91 @@ def trunc(a: A) -> A:
587
587
 
588
588
 
589
589
  # Reduction functions
590
- 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:
591
591
  """
592
592
  Named version of [jax.numpy.all](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.all.html#jax.numpy.all).
593
593
  """
594
594
  return wrap_reduction_call(jnp.all, array, axis, where, single_axis_only=False, supports_where=True)
595
595
 
596
596
 
597
- 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:
598
598
  """
599
599
  Aliax for max. See max for details.
600
600
  """
601
601
  return wrap_reduction_call(jnp.amax, array, axis, where, single_axis_only=False, supports_where=True)
602
602
 
603
603
 
604
- 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:
605
605
  """
606
606
  Aliax for min. See min for details.
607
607
  """
608
608
  return wrap_reduction_call(jnp.amin, array, axis, where, single_axis_only=False, supports_where=True)
609
609
 
610
610
 
611
- 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:
612
612
  """True if any elements along a given axis or axes are True. If axis is None, any elements are True."""
613
613
  return wrap_reduction_call(jnp.any, array, axis, where, single_axis_only=False, supports_where=True)
614
614
 
615
615
 
616
- def argmax(array: NamedArray, axis: Optional[AxisSelector]) -> NamedArray:
616
+ def argmax(array: NamedArray, axis: AxisSelector | None) -> NamedArray:
617
617
  return wrap_reduction_call(jnp.argmax, array, axis, None, single_axis_only=True, supports_where=False)
618
618
 
619
619
 
620
- def argmin(array: NamedArray, axis: Optional[AxisSelector]) -> NamedArray:
620
+ def argmin(array: NamedArray, axis: AxisSelector | None) -> NamedArray:
621
621
  return wrap_reduction_call(jnp.argmin, array, axis, None, single_axis_only=True, supports_where=False)
622
622
 
623
623
 
624
- 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:
625
625
  return wrap_reduction_call(jnp.max, array, axis, where, single_axis_only=False, supports_where=True)
626
626
 
627
627
 
628
628
  def mean(
629
629
  array: NamedArray,
630
- axis: Optional[AxisSelection] = None,
630
+ axis: AxisSelection | None = None,
631
631
  *,
632
- where: Optional[NamedArray] = None,
633
- dtype: Optional[DTypeLike] = None,
632
+ where: NamedArray | None = None,
633
+ dtype: DTypeLike | None = None,
634
634
  ) -> NamedArray:
635
635
  return wrap_reduction_call(jnp.mean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
636
636
 
637
637
 
638
- 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:
639
639
  return wrap_reduction_call(jnp.min, array, axis, where, single_axis_only=False, supports_where=True)
640
640
 
641
641
 
642
642
  def prod(
643
643
  array: NamedArray,
644
- axis: Optional[AxisSelection] = None,
644
+ axis: AxisSelection | None = None,
645
645
  *,
646
- where: Optional[NamedArray] = None,
647
- dtype: Optional[DTypeLike] = None,
646
+ where: NamedArray | None = None,
647
+ dtype: DTypeLike | None = None,
648
648
  ) -> NamedArray:
649
649
  return wrap_reduction_call(jnp.prod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
650
650
 
651
651
 
652
652
  def std(
653
653
  array: NamedArray,
654
- axis: Optional[AxisSelection] = None,
654
+ axis: AxisSelection | None = None,
655
655
  *,
656
- where: Optional[NamedArray] = None,
656
+ where: NamedArray | None = None,
657
657
  ddof: int = 0,
658
- dtype: Optional[DTypeLike] = None,
658
+ dtype: DTypeLike | None = None,
659
659
  ) -> NamedArray:
660
660
  return wrap_reduction_call(
661
661
  jnp.std, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
662
662
  )
663
663
 
664
664
 
665
- 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:
666
666
  return wrap_reduction_call(jnp.ptp, array, axis, where, single_axis_only=False, supports_where=True)
667
667
 
668
668
 
669
669
  def product(
670
670
  array: NamedArray,
671
- axis: Optional[AxisSelection] = None,
671
+ axis: AxisSelection | None = None,
672
672
  *,
673
- where: Optional[NamedArray] = None,
674
- dtype: Optional[DTypeLike] = None,
673
+ where: NamedArray | None = None,
674
+ dtype: DTypeLike | None = None,
675
675
  ) -> NamedArray:
676
676
  return wrap_reduction_call(
677
677
  jnp.product, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
@@ -683,50 +683,50 @@ _sum = sum
683
683
 
684
684
  def sum(
685
685
  array: NamedArray,
686
- axis: Optional[AxisSelection] = None,
686
+ axis: AxisSelection | None = None,
687
687
  *,
688
- where: Optional[NamedArray] = None,
689
- dtype: Optional[DTypeLike] = None,
688
+ where: NamedArray | None = None,
689
+ dtype: DTypeLike | None = None,
690
690
  ) -> NamedArray:
691
691
  return wrap_reduction_call(jnp.sum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype)
692
692
 
693
693
 
694
694
  def var(
695
695
  array: NamedArray,
696
- axis: Optional[AxisSelection] = None,
696
+ axis: AxisSelection | None = None,
697
697
  *,
698
- where: Optional[NamedArray] = None,
698
+ where: NamedArray | None = None,
699
699
  ddof: int = 0,
700
- dtype: Optional[DTypeLike] = None,
700
+ dtype: DTypeLike | None = None,
701
701
  ) -> NamedArray:
702
702
  return wrap_reduction_call(
703
703
  jnp.var, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
704
704
  )
705
705
 
706
706
 
707
- def nanargmax(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
707
+ def nanargmax(array: NamedArray, axis: AxisSelector | None = None) -> NamedArray:
708
708
  return wrap_reduction_call(jnp.nanargmax, array, axis, None, single_axis_only=True, supports_where=False)
709
709
 
710
710
 
711
- def nanargmin(array: NamedArray, axis: Optional[AxisSelector] = None) -> NamedArray:
711
+ def nanargmin(array: NamedArray, axis: AxisSelector | None = None) -> NamedArray:
712
712
  return wrap_reduction_call(jnp.nanargmin, array, axis, None, single_axis_only=True, supports_where=False)
713
713
 
714
714
 
715
715
  def nanmax(
716
716
  array: NamedArray,
717
- axis: Optional[AxisSelection] = None,
717
+ axis: AxisSelection | None = None,
718
718
  *,
719
- where: Optional[NamedArray] = None,
719
+ where: NamedArray | None = None,
720
720
  ) -> NamedArray:
721
721
  return wrap_reduction_call(jnp.nanmax, array, axis, where, single_axis_only=False, supports_where=True)
722
722
 
723
723
 
724
724
  def nanmean(
725
725
  array: NamedArray,
726
- axis: Optional[AxisSelection] = None,
726
+ axis: AxisSelection | None = None,
727
727
  *,
728
- where: Optional[NamedArray] = None,
729
- dtype: Optional[DTypeLike] = None,
728
+ where: NamedArray | None = None,
729
+ dtype: DTypeLike | None = None,
730
730
  ) -> NamedArray:
731
731
  return wrap_reduction_call(
732
732
  jnp.nanmean, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
@@ -735,19 +735,19 @@ def nanmean(
735
735
 
736
736
  def nanmin(
737
737
  array: NamedArray,
738
- axis: Optional[AxisSelection] = None,
738
+ axis: AxisSelection | None = None,
739
739
  *,
740
- where: Optional[NamedArray] = None,
740
+ where: NamedArray | None = None,
741
741
  ) -> NamedArray:
742
742
  return wrap_reduction_call(jnp.nanmin, array, axis, where, single_axis_only=False, supports_where=True)
743
743
 
744
744
 
745
745
  def nanprod(
746
746
  array: NamedArray,
747
- axis: Optional[AxisSelection] = None,
747
+ axis: AxisSelection | None = None,
748
748
  *,
749
- where: Optional[NamedArray] = None,
750
- dtype: Optional[DTypeLike] = None,
749
+ where: NamedArray | None = None,
750
+ dtype: DTypeLike | None = None,
751
751
  ) -> NamedArray:
752
752
  return wrap_reduction_call(
753
753
  jnp.nanprod, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
@@ -756,11 +756,11 @@ def nanprod(
756
756
 
757
757
  def nanstd(
758
758
  array: NamedArray,
759
- axis: Optional[AxisSelection] = None,
759
+ axis: AxisSelection | None = None,
760
760
  *,
761
- where: Optional[NamedArray] = None,
761
+ where: NamedArray | None = None,
762
762
  ddof: int = 0,
763
- dtype: Optional[DTypeLike] = None,
763
+ dtype: DTypeLike | None = None,
764
764
  ) -> NamedArray:
765
765
  return wrap_reduction_call(
766
766
  jnp.nanstd, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
@@ -769,10 +769,10 @@ def nanstd(
769
769
 
770
770
  def nansum(
771
771
  array: NamedArray,
772
- axis: Optional[AxisSelection] = None,
772
+ axis: AxisSelection | None = None,
773
773
  *,
774
- where: Optional[NamedArray] = None,
775
- dtype: Optional[DTypeLike] = None,
774
+ where: NamedArray | None = None,
775
+ dtype: DTypeLike | None = None,
776
776
  ) -> NamedArray:
777
777
  return wrap_reduction_call(
778
778
  jnp.nansum, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype
@@ -781,11 +781,11 @@ def nansum(
781
781
 
782
782
  def nanvar(
783
783
  array: NamedArray,
784
- axis: Optional[AxisSelection] = None,
784
+ axis: AxisSelection | None = None,
785
785
  *,
786
- where: Optional[NamedArray] = None,
786
+ where: NamedArray | None = None,
787
787
  ddof: int = 0,
788
- dtype: Optional[DTypeLike] = None,
788
+ dtype: DTypeLike | None = None,
789
789
  ) -> NamedArray:
790
790
  return wrap_reduction_call(
791
791
  jnp.nanvar, array, axis, where, single_axis_only=False, supports_where=True, dtype=dtype, ddof=ddof
@@ -795,28 +795,28 @@ def nanvar(
795
795
  # "Normalization" functions that use an axis but don't change the shape
796
796
 
797
797
 
798
- def cumsum(a: NamedArray, axis: AxisSelector, *, dtype: Optional[DTypeLike] = None) -> NamedArray:
798
+ def cumsum(a: NamedArray, axis: AxisSelector, *, dtype: DTypeLike | None = None) -> NamedArray:
799
799
  """
800
800
  Named version of [jax.numpy.cumsum](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.cumsum.html)
801
801
  """
802
802
  return wrap_axiswise_call(jnp.cumsum, a, axis, dtype=dtype, single_axis_only=True)
803
803
 
804
804
 
805
- def cumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None) -> NamedArray:
805
+ def cumprod(a: NamedArray, axis: AxisSelector, dtype: DTypeLike | None = None) -> NamedArray:
806
806
  """
807
807
  Named version of [jax.numpy.cumprod](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.cumprod.html)
808
808
  """
809
809
  return wrap_axiswise_call(jnp.cumprod, a, axis, dtype=dtype, single_axis_only=True)
810
810
 
811
811
 
812
- def nancumsum(a: NamedArray, axis: AxisSelector, *, dtype: Optional[DTypeLike] = None) -> NamedArray:
812
+ def nancumsum(a: NamedArray, axis: AxisSelector, *, dtype: DTypeLike | None = None) -> NamedArray:
813
813
  """
814
814
  Named version of [jax.numpy.nancumsum](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumsum.html)
815
815
  """
816
816
  return wrap_axiswise_call(jnp.nancumsum, a, axis, dtype=dtype, single_axis_only=True)
817
817
 
818
818
 
819
- def nancumprod(a: NamedArray, axis: AxisSelector, dtype: Optional[DTypeLike] = None) -> NamedArray:
819
+ def nancumprod(a: NamedArray, axis: AxisSelector, dtype: DTypeLike | None = None) -> NamedArray:
820
820
  """
821
821
  Named version of [jax.numpy.nancumprod](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.nancumprod.html)
822
822
  """
@@ -6,7 +6,6 @@
6
6
  import functools as ft
7
7
  import typing
8
8
  import warnings
9
- from typing import Dict, Optional, Tuple
10
9
 
11
10
  import jax
12
11
 
@@ -29,11 +28,11 @@ from haliax.types import DTypeLike, PrecisionLike
29
28
  # deprecated overload
30
29
  @typing.overload
31
30
  def dot(
32
- axis: Optional[AxisSelection],
31
+ axis: AxisSelection | None,
33
32
  *arrays: NamedArray,
34
33
  precision: PrecisionLike = None,
35
- preferred_element_type: Optional[DTypeLike] = None,
36
- out_axes: Optional[PartialAxisSpec] = ...,
34
+ preferred_element_type: DTypeLike | None = None,
35
+ out_axes: PartialAxisSpec | None = ...,
37
36
  dot_general=jax.lax.dot_general,
38
37
  ) -> NamedArray: ...
39
38
 
@@ -41,10 +40,10 @@ def dot(
41
40
  @typing.overload
42
41
  def dot(
43
42
  *arrays: NamedArray,
44
- axis: Optional[AxisSelection],
43
+ axis: AxisSelection | None,
45
44
  precision: PrecisionLike = None,
46
- preferred_element_type: Optional[DTypeLike] = None,
47
- out_axes: Optional[PartialAxisSpec] = ...,
45
+ preferred_element_type: DTypeLike | None = None,
46
+ out_axes: PartialAxisSpec | None = ...,
48
47
  dot_general=jax.lax.dot_general,
49
48
  ) -> NamedArray: ...
50
49
 
@@ -52,8 +51,8 @@ def dot(
52
51
  def dot(
53
52
  *arrays,
54
53
  precision: PrecisionLike = None,
55
- preferred_element_type: Optional[DTypeLike] = None,
56
- out_axes: Optional[PartialAxisSpec] = None,
54
+ preferred_element_type: DTypeLike | None = None,
55
+ out_axes: PartialAxisSpec | None = None,
57
56
  dot_general=jax.lax.dot_general,
58
57
  **kwargs,
59
58
  ) -> NamedArray:
@@ -82,7 +81,7 @@ def dot(
82
81
  which in turn passes it to jax.lax.dot_general.
83
82
  preferred_element_type (DTypeLike, optional): The preferred element type of the result. Defaults to None.
84
83
  This argument is passed to `jax.numpy.einsum`.
85
- out_axes (Optional[PartialAxisSpec], optional): a potentially partial specification of the output axes.
84
+ out_axes (PartialAxisSpec | None, optional): a potentially partial specification of the output axes.
86
85
  If provided, the output will be transposed to match the provided axes. Defaults to None.
87
86
 
88
87
 
@@ -107,8 +106,8 @@ def dot(
107
106
  # to call dot_general we need two things:
108
107
  # list of contractions and list of arrays
109
108
 
110
- all_axes: Tuple[Axis, ...] = ft.reduce(union_axes, (a.axes for a in arrays), ()) # type: ignore
111
- output_axes: Tuple[Axis, ...]
109
+ all_axes: tuple[Axis, ...] = ft.reduce(union_axes, (a.axes for a in arrays), ()) # type: ignore
110
+ output_axes: tuple[Axis, ...]
112
111
  if axis is None:
113
112
  # we want to contract over all the axes
114
113
  output_axes = ()
@@ -121,7 +120,7 @@ def dot(
121
120
  array_specs = []
122
121
 
123
122
  next_index = 0
124
- axis_mappings: Dict[str, int] = {}
123
+ axis_mappings: dict[str, int] = {}
125
124
 
126
125
  for a in arrays:
127
126
  spec = ""
@@ -5,7 +5,6 @@
5
5
 
6
6
  import functools
7
7
  from types import EllipsisType
8
- from typing import Optional, Tuple
9
8
 
10
9
  import jax.lax
11
10
 
@@ -24,7 +23,7 @@ def einsum(
24
23
  equation: str,
25
24
  *arrays: NamedArray,
26
25
  precision: PrecisionLike = None,
27
- preferred_element_type: Optional[DTypeLike] = None,
26
+ preferred_element_type: DTypeLike | None = None,
28
27
  _dot_general: DotGeneralOp = jax.lax.dot_general,
29
28
  **axis_aliases: AxisSelector,
30
29
  ) -> NamedArray:
@@ -306,7 +305,7 @@ def _all_input_axes(arrays):
306
305
  return ensure_tuple(functools.reduce(union_axes, (a.axes for a in arrays), ())) # type: ignore
307
306
 
308
307
 
309
- def _captures_to_axis_names(equation, lhs, aliases) -> Tuple[list[str | EllipsisType], bool, set[str]]:
308
+ def _captures_to_axis_names(equation, lhs, aliases) -> tuple[list[str | EllipsisType], bool, set[str]]:
310
309
  covered_aliases = set()
311
310
  candidate_axes: list[str | EllipsisType] = []
312
311
  has_ellipsis = False
@@ -5,16 +5,16 @@
5
5
 
6
6
  import dataclasses
7
7
  from types import EllipsisType
8
- from typing import Mapping, NoReturn, Optional, Sequence
8
+ from typing import Mapping, NoReturn, Sequence
9
9
 
10
10
  from haliax.axis import Axis, AxisSelector
11
11
 
12
12
 
13
13
  @dataclasses.dataclass(frozen=True)
14
14
  class _AxisCapture:
15
- binding: Optional[str] = None
15
+ binding: str | None = None
16
16
  axes: tuple[str, ...] = ()
17
- char_range: Optional[tuple[int, int]] = None
17
+ char_range: tuple[int, int] | None = None
18
18
 
19
19
  def __post_init__(self):
20
20
  if len(self.axes) == 0:
@@ -27,7 +27,7 @@ class Expression:
27
27
  is_ordered: bool
28
28
 
29
29
 
30
- def raise_parse_error(message: str, expression: str, pos: Optional[int | tuple[int, int]]) -> NoReturn:
30
+ def raise_parse_error(message: str, expression: str, pos: int | tuple[int, int] | None) -> NoReturn:
31
31
  """Raise a ValueError with a message and the position in the expression."""
32
32
  fmt = f"Error while parsing:\n {expression}"
33
33
  if pos is not None:
@@ -234,7 +234,7 @@ class AliasTable:
234
234
  else:
235
235
  self.bindings = {**bindings}
236
236
 
237
- def dealias_binding(self, binding: str) -> Optional[AxisSelector]:
237
+ def dealias_binding(self, binding: str) -> AxisSelector | None:
238
238
  return self.bindings.get(binding, None)
239
239
 
240
240
  def bind_alias(self, alias: str, axis: Axis, expr, char_range):
@@ -7,7 +7,7 @@
7
7
  import dataclasses
8
8
  import typing
9
9
  from types import EllipsisType
10
- from typing import Mapping, Optional, Sequence
10
+ from typing import Mapping, Sequence
11
11
 
12
12
  import jax.lax
13
13
  import jax.numpy as jnp
@@ -157,7 +157,7 @@ def einops_rearrange(array: NamedArray, expression: str, **bindings: AxisSelecto
157
157
  @dataclasses.dataclass(frozen=True)
158
158
  class _Plan:
159
159
  intermediate_axes: tuple[Axis, ...]
160
- transpose: Optional[tuple[int, ...]]
160
+ transpose: tuple[int, ...] | None
161
161
  needs_final_reshape: bool
162
162
 
163
163
  final_axes: tuple[Axis, ...]
@@ -170,7 +170,7 @@ def _plan_rearrange(
170
170
  grouped_new_shapes = _determine_initial_reshape(original_str, lhs, array, aliases)
171
171
  intermediate_axes = tuple(ax for split_axes in grouped_new_shapes for ax in split_axes)
172
172
 
173
- transpose: Optional[tuple[int, ...]]
173
+ transpose: tuple[int, ...] | None
174
174
  transpose, final_axes = _determine_final_transpose_and_reshape(original_str, rhs, aliases, intermediate_axes)
175
175
 
176
176
  transposed_intermediate_axes = tuple(intermediate_axes[i] for i in transpose)
@@ -289,7 +289,7 @@ def _determine_initial_reshape(
289
289
  # the lhs all need to be bound to axes in the array, or synthesized as parts of axes.
290
290
  # In the lhs, bindings look like either a name, or a name and a list of (new) axes.
291
291
  # bindings can either be done by name, or by position, depending on if lhs.is_ordered
292
- new_shapes: list[Optional[list[Axis]]] = [None] * len(array.axes)
292
+ new_shapes: list[list[Axis] | None] = [None] * len(array.axes)
293
293
  used_new_names: set[str] = set() # names can only be used once on a side
294
294
 
295
295
  # one subtle difference between the lhs and the rhs is the handling of binding in expressions like (a: b c)
@@ -301,7 +301,7 @@ def _determine_initial_reshape(
301
301
  # if we start with an ellipsis, we bind from the right
302
302
  # if we end with an ellipsis, we bind from the left
303
303
  ellipsis_pos = None
304
- axis_index_for_capture: list[Optional[int]] = [None] * len(lhs.captures)
304
+ axis_index_for_capture: list[int | None] = [None] * len(lhs.captures)
305
305
  covered_axes = set()
306
306
  # bind from the left
307
307
  axis_pos = 0
@@ -414,8 +414,8 @@ def _solve_split_axes(axis, capture, aliases, used_new_names, expression):
414
414
  """
415
415
  Given an axis and a capture of the form (a: b c) or (b c) on the lhs, solve for the new axes.
416
416
  """
417
- new_axes: list[Optional[Axis]] = []
418
- unsolved_axis_index: Optional[int] = None
417
+ new_axes: list[Axis | None] = []
418
+ unsolved_axis_index: int | None = None
419
419
 
420
420
  # easy case: 1 axis in capture
421
421
  if len(capture.axes) == 1: