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.
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/publish_dev.yaml +1 -1
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_pre_commit.yaml +1 -1
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_quick_levanter_tests.yaml +1 -1
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.github/workflows/run_tests.yaml +2 -2
- {haliax-1.4.dev419 → haliax-1.4.dev438}/PKG-INFO +2 -2
- {haliax-1.4.dev419 → haliax-1.4.dev438}/pyproject.toml +2 -2
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/__about__.py +1 -1
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/__init__.py +62 -62
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/dot.py +12 -13
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/einsum.py +2 -3
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/parsing.py +5 -5
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/rearrange.py +7 -7
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/scan.py +8 -7
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/state_dict.py +15 -15
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/axis.py +14 -14
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/core.py +95 -98
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/debug.py +5 -5
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/jax_utils.py +8 -8
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/attention.py +14 -15
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/conv.py +4 -4
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/dropout.py +4 -6
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/embedding.py +24 -11
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/linear.py +56 -12
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/loss.py +20 -21
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/mlp.py +2 -2
- haliax-1.4.dev438/src/haliax/nn/mup.py +206 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/normalization.py +14 -14
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/pool.py +5 -5
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/scan.py +9 -11
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/ops.py +6 -6
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/partitioning.py +44 -42
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/quantization.py +4 -4
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/random.py +2 -5
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/specialized_fns.py +3 -5
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/state_dict.py +2 -2
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/tree_util.py +3 -2
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/types.py +11 -11
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/util.py +3 -3
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/wrap.py +7 -7
- haliax-1.4.dev438/tests/test_mup_coordinate_check.py +164 -0
- haliax-1.4.dev438/tests/test_mup_embedding.py +48 -0
- haliax-1.4.dev438/tests/test_mup_linear.py +120 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/uv.lock +43 -439
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.agents/projects/api_parity.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.coveragerc +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.flake8 +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.gitignore +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/AGENTS.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/AUTHORS.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/CONTRIBUTORS.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/LICENSE +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/README.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/api.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/css/material.css +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/faq.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/fp8.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/index.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/indexing.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/matmul.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/nn.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/partitioning.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/primer.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/rearrange.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/requirements.txt +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/scan.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/state-dict.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/tutorial.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/typing.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/docs/vmap.md +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/etc/license_header.txt +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/mkdocs.yml +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/fft.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/field.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/src/haliax/poly.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/core_test.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_attention.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_axis.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_bitwise_ops.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_conv.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_debug.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_dot.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_fft.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_field.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_hof.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_int8.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_nan_reductions.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_nn.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_ops.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_poly_ops.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_pool.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_random.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_scan.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_utils.py +0 -0
- {haliax-1.4.dev419 → haliax-1.4.dev438}/tests/test_visualize_sharding.py +0 -0
|
@@ -9,10 +9,10 @@ jobs:
|
|
|
9
9
|
|
|
10
10
|
steps:
|
|
11
11
|
- uses: actions/checkout@v3
|
|
12
|
-
- name: Set up Python 3.
|
|
12
|
+
- name: Set up Python 3.11
|
|
13
13
|
uses: actions/setup-python@v4
|
|
14
14
|
with:
|
|
15
|
-
python-version: 3.
|
|
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.
|
|
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.
|
|
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.
|
|
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 = ["
|
|
63
|
+
target-version = ["py311"]
|
|
64
64
|
preview = true
|
|
65
65
|
|
|
66
66
|
[tool.isort]
|
|
@@ -4,7 +4,7 @@
|
|
|
4
4
|
|
|
5
5
|
|
|
6
6
|
import typing as t
|
|
7
|
-
from typing import
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
630
|
+
axis: AxisSelection | None = None,
|
|
631
631
|
*,
|
|
632
|
-
where:
|
|
633
|
-
dtype:
|
|
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:
|
|
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:
|
|
644
|
+
axis: AxisSelection | None = None,
|
|
645
645
|
*,
|
|
646
|
-
where:
|
|
647
|
-
dtype:
|
|
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:
|
|
654
|
+
axis: AxisSelection | None = None,
|
|
655
655
|
*,
|
|
656
|
-
where:
|
|
656
|
+
where: NamedArray | None = None,
|
|
657
657
|
ddof: int = 0,
|
|
658
|
-
dtype:
|
|
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:
|
|
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:
|
|
671
|
+
axis: AxisSelection | None = None,
|
|
672
672
|
*,
|
|
673
|
-
where:
|
|
674
|
-
dtype:
|
|
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:
|
|
686
|
+
axis: AxisSelection | None = None,
|
|
687
687
|
*,
|
|
688
|
-
where:
|
|
689
|
-
dtype:
|
|
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:
|
|
696
|
+
axis: AxisSelection | None = None,
|
|
697
697
|
*,
|
|
698
|
-
where:
|
|
698
|
+
where: NamedArray | None = None,
|
|
699
699
|
ddof: int = 0,
|
|
700
|
-
dtype:
|
|
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:
|
|
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:
|
|
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:
|
|
717
|
+
axis: AxisSelection | None = None,
|
|
718
718
|
*,
|
|
719
|
-
where:
|
|
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:
|
|
726
|
+
axis: AxisSelection | None = None,
|
|
727
727
|
*,
|
|
728
|
-
where:
|
|
729
|
-
dtype:
|
|
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:
|
|
738
|
+
axis: AxisSelection | None = None,
|
|
739
739
|
*,
|
|
740
|
-
where:
|
|
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:
|
|
747
|
+
axis: AxisSelection | None = None,
|
|
748
748
|
*,
|
|
749
|
-
where:
|
|
750
|
-
dtype:
|
|
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:
|
|
759
|
+
axis: AxisSelection | None = None,
|
|
760
760
|
*,
|
|
761
|
-
where:
|
|
761
|
+
where: NamedArray | None = None,
|
|
762
762
|
ddof: int = 0,
|
|
763
|
-
dtype:
|
|
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:
|
|
772
|
+
axis: AxisSelection | None = None,
|
|
773
773
|
*,
|
|
774
|
-
where:
|
|
775
|
-
dtype:
|
|
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:
|
|
784
|
+
axis: AxisSelection | None = None,
|
|
785
785
|
*,
|
|
786
|
-
where:
|
|
786
|
+
where: NamedArray | None = None,
|
|
787
787
|
ddof: int = 0,
|
|
788
|
-
dtype:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
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:
|
|
31
|
+
axis: AxisSelection | None,
|
|
33
32
|
*arrays: NamedArray,
|
|
34
33
|
precision: PrecisionLike = None,
|
|
35
|
-
preferred_element_type:
|
|
36
|
-
out_axes:
|
|
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:
|
|
43
|
+
axis: AxisSelection | None,
|
|
45
44
|
precision: PrecisionLike = None,
|
|
46
|
-
preferred_element_type:
|
|
47
|
-
out_axes:
|
|
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:
|
|
56
|
-
out_axes:
|
|
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 (
|
|
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:
|
|
111
|
-
output_axes:
|
|
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:
|
|
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:
|
|
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) ->
|
|
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,
|
|
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:
|
|
15
|
+
binding: str | None = None
|
|
16
16
|
axes: tuple[str, ...] = ()
|
|
17
|
-
char_range:
|
|
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:
|
|
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) ->
|
|
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,
|
|
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:
|
|
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:
|
|
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[
|
|
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[
|
|
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[
|
|
418
|
-
unsolved_axis_index:
|
|
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:
|