haliax 1.4.dev408__tar.gz → 1.4.dev410__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.dev408 → haliax-1.4.dev410}/.agents/projects/api_parity.md +6 -6
- haliax-1.4.dev410/.pre-commit-config.yaml +43 -0
- haliax-1.4.dev410/AUTHORS.md +5 -0
- haliax-1.4.dev410/CONTRIBUTORS.md +15 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/PKG-INFO +2 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/api.md +6 -0
- haliax-1.4.dev410/etc/license_header.txt +3 -0
- haliax-1.4.dev410/src/haliax/__about__.py +6 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/__init__.py +41 -3
- haliax-1.4.dev410/src/haliax/_src/__init__.py +3 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/compile_utils.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/dot.py +7 -4
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/einsum.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/fp8.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/parsing.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/rearrange.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/scan.py +11 -14
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/state_dict.py +9 -7
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/_src/util.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/axis.py +24 -38
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/core.py +25 -41
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/debug.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/field.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/haxtyping.py +49 -23
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/hof.py +10 -3
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/jax_utils.py +7 -5
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/__init__.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/activations.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/attention.py +6 -2
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/conv.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/dropout.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/embedding.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/linear.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/loss.py +9 -8
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/mlp.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/normalization.py +6 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/pool.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/nn/scan.py +69 -62
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/ops.py +50 -14
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/partitioning.py +13 -12
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/quantization.py +6 -3
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/random.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/specialized_fns.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/state_dict.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/tree_util.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/types.py +6 -3
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/util.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/src/haliax/wrap.py +7 -4
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/core_test.py +6 -3
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_attention.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_axis.py +5 -0
- haliax-1.4.dev410/tests/test_bitwise_ops.py +45 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_conv.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_debug.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_dot.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_dtype_typing.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_einsum.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_field.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_fp8.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_hof.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_int8.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_moe_linear.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_namedarray_typing.py +6 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_nn.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_ops.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_parsing.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_partitioning.py +9 -3
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_pool.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_random.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_rearrange.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_scan.py +6 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_scatter_gather.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_specialized_fns.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_state_dict.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_tree_util.py +5 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_utils.py +5 -1
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_visualize_sharding.py +5 -0
- haliax-1.4.dev408/.pre-commit-config.yaml +0 -39
- haliax-1.4.dev408/src/haliax/__about__.py +0 -1
- haliax-1.4.dev408/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.coveragerc +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.flake8 +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.gitignore +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/AGENTS.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/LICENSE +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/README.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/css/material.css +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/faq.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/fp8.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/index.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/indexing.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/matmul.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/nn.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/partitioning.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/primer.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/rearrange.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/requirements.txt +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/scan.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/state-dict.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/tutorial.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/typing.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/docs/vmap.md +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/mkdocs.yml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/pyproject.toml +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/tests/test_nan_reductions.py +0 -0
- {haliax-1.4.dev408 → haliax-1.4.dev410}/uv.lock +0 -0
|
@@ -18,10 +18,10 @@ APIs that don't translate well to named tensors are intentionally omitted here.
|
|
|
18
18
|
- [ ] `atan2`
|
|
19
19
|
- [ ] `average`
|
|
20
20
|
- [ ] `bartlett`
|
|
21
|
-
- [
|
|
22
|
-
- [
|
|
23
|
-
- [
|
|
24
|
-
- [
|
|
21
|
+
- [x] `bitwise_count`
|
|
22
|
+
- [x] `bitwise_invert`
|
|
23
|
+
- [x] `bitwise_left_shift`
|
|
24
|
+
- [x] `bitwise_right_shift`
|
|
25
25
|
- [ ] `blackman`
|
|
26
26
|
- [ ] `block`
|
|
27
27
|
- [ ] `broadcast_shapes`
|
|
@@ -109,7 +109,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
|
|
|
109
109
|
- [x] `nanvar`
|
|
110
110
|
- [ ] `nonzero`
|
|
111
111
|
- [ ] `ogrid`
|
|
112
|
-
- [
|
|
112
|
+
- [x] `packbits`
|
|
113
113
|
- [ ] `partition`
|
|
114
114
|
- [ ] `percentile`
|
|
115
115
|
- [ ] `permute_dims`
|
|
@@ -152,7 +152,7 @@ APIs that don't translate well to named tensors are intentionally omitted here.
|
|
|
152
152
|
- [ ] `triu_indices`
|
|
153
153
|
- [ ] `triu_indices_from`
|
|
154
154
|
- [ ] `union1d`
|
|
155
|
-
- [
|
|
155
|
+
- [x] `unpackbits`
|
|
156
156
|
- [ ] `unravel_index`
|
|
157
157
|
- [ ] `unstack`
|
|
158
158
|
- [ ] `unwrap`
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
# See https://pre-commit.com
|
|
2
|
+
exclude: ".git|.venv|tests/snapshots/.*/.*"
|
|
3
|
+
|
|
4
|
+
repos:
|
|
5
|
+
- repo: https://github.com/astral-sh/ruff-pre-commit
|
|
6
|
+
rev: v0.11.10
|
|
7
|
+
hooks:
|
|
8
|
+
- id: ruff
|
|
9
|
+
args: [ --fix, --exit-non-zero-on-fix ]
|
|
10
|
+
|
|
11
|
+
- repo: https://github.com/Lucas-C/pre-commit-hooks
|
|
12
|
+
rev: v1.5.5
|
|
13
|
+
hooks:
|
|
14
|
+
- id: insert-license
|
|
15
|
+
files: \.py$
|
|
16
|
+
args:
|
|
17
|
+
- --license-filepath
|
|
18
|
+
- etc/license_header.txt
|
|
19
|
+
- --use-current-year
|
|
20
|
+
|
|
21
|
+
- repo: https://github.com/psf/black
|
|
22
|
+
rev: 25.1.0
|
|
23
|
+
hooks:
|
|
24
|
+
- id: black
|
|
25
|
+
|
|
26
|
+
- repo: https://github.com/pre-commit/pre-commit-hooks
|
|
27
|
+
rev: v5.0.0
|
|
28
|
+
hooks:
|
|
29
|
+
- id: check-added-large-files
|
|
30
|
+
- id: check-ast
|
|
31
|
+
- id: check-case-conflict
|
|
32
|
+
- id: check-merge-conflict
|
|
33
|
+
- id: check-toml
|
|
34
|
+
- id: check-yaml
|
|
35
|
+
args: [ --unsafe ]
|
|
36
|
+
- id: end-of-file-fixer
|
|
37
|
+
- id: trailing-whitespace
|
|
38
|
+
|
|
39
|
+
- repo: https://github.com/pre-commit/mirrors-mypy
|
|
40
|
+
rev: v1.16.1
|
|
41
|
+
hooks:
|
|
42
|
+
- id: mypy
|
|
43
|
+
args: [--ignore-missing-imports, --check-untyped-defs]
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
# Contributors
|
|
2
|
+
|
|
3
|
+
The following individuals have contributed to Haliax:
|
|
4
|
+
|
|
5
|
+
- David Hall <dlwh@stanford.edu>
|
|
6
|
+
- David Hall <dlwh@cs.stanford.edu>
|
|
7
|
+
- Jason Wang <blahblahj.wsy@gmail.com>
|
|
8
|
+
- Ivan Zhou <ivan.zhouyq@gmail.com>
|
|
9
|
+
- rohan-mehta-1024 <69774557+rohan-mehta-1024@users.noreply.github.com>
|
|
10
|
+
- Gary Miguel <garymm@garymm.org>
|
|
11
|
+
- Jennifer Zhou <jennifer@jezh.me>
|
|
12
|
+
- Joseph Camacho <camacho.joseph@gmail.com>
|
|
13
|
+
- Omead Pooladzandi <opooladz@ucla.edu>
|
|
14
|
+
- Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com>
|
|
15
|
+
- Russell Power <russell.power@gmail.com>
|
|
@@ -1,11 +1,12 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev410
|
|
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/
|
|
7
7
|
Project-URL: Documentation, https://haliax.readthedocs.io/en/latest/
|
|
8
8
|
Author-email: David Hall <dlwh@cs.stanford.edu>
|
|
9
|
+
License-File: AUTHORS.md
|
|
9
10
|
License-File: LICENSE
|
|
10
11
|
Classifier: Development Status :: 4 - Beta
|
|
11
12
|
Classifier: Intended Audience :: Science/Research
|
|
@@ -176,6 +176,8 @@ These are all more or less directly from JAX's NumPy API.
|
|
|
176
176
|
::: haliax.arctan
|
|
177
177
|
::: haliax.arctanh
|
|
178
178
|
::: haliax.around
|
|
179
|
+
::: haliax.bitwise_count
|
|
180
|
+
::: haliax.bitwise_invert
|
|
179
181
|
::: haliax.bitwise_not
|
|
180
182
|
::: haliax.cbrt
|
|
181
183
|
::: haliax.ceil
|
|
@@ -233,7 +235,9 @@ These are all more or less directly from JAX's NumPy API.
|
|
|
233
235
|
::: haliax.add
|
|
234
236
|
::: haliax.arctan2
|
|
235
237
|
::: haliax.bitwise_and
|
|
238
|
+
::: haliax.bitwise_left_shift
|
|
236
239
|
::: haliax.bitwise_or
|
|
240
|
+
::: haliax.bitwise_right_shift
|
|
237
241
|
::: haliax.bitwise_xor
|
|
238
242
|
::: haliax.divide
|
|
239
243
|
::: haliax.divmod
|
|
@@ -270,6 +274,8 @@ These are all more or less directly from JAX's NumPy API.
|
|
|
270
274
|
|
|
271
275
|
::: haliax.bincount
|
|
272
276
|
::: haliax.clip
|
|
277
|
+
::: haliax.packbits
|
|
278
|
+
::: haliax.unpackbits
|
|
273
279
|
::: haliax.isclose
|
|
274
280
|
::: haliax.allclose
|
|
275
281
|
::: haliax.array_equal
|
|
@@ -1,10 +1,14 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
import typing as t
|
|
2
7
|
from typing import Optional, Sequence
|
|
3
8
|
|
|
4
9
|
import jax
|
|
5
10
|
import jax.numpy as jnp
|
|
6
11
|
|
|
7
|
-
|
|
8
12
|
try:
|
|
9
13
|
from jax.typing import DTypeLike
|
|
10
14
|
except ImportError:
|
|
@@ -44,7 +48,9 @@ from .axis import (
|
|
|
44
48
|
)
|
|
45
49
|
from .core import (
|
|
46
50
|
NamedArray,
|
|
47
|
-
NamedArrayAxes,
|
|
51
|
+
NamedArrayAxes,
|
|
52
|
+
NamedArrayAxesSpec,
|
|
53
|
+
NamedOrNumeric,
|
|
48
54
|
are_shape_checks_enabled,
|
|
49
55
|
broadcast_arrays,
|
|
50
56
|
broadcast_axis,
|
|
@@ -83,6 +89,8 @@ from .ops import (
|
|
|
83
89
|
unique_counts,
|
|
84
90
|
unique_inverse,
|
|
85
91
|
unique_all,
|
|
92
|
+
packbits,
|
|
93
|
+
unpackbits,
|
|
86
94
|
searchsorted,
|
|
87
95
|
bincount,
|
|
88
96
|
where,
|
|
@@ -100,7 +108,6 @@ from .wrap import (
|
|
|
100
108
|
wrap_reduction_call,
|
|
101
109
|
)
|
|
102
110
|
|
|
103
|
-
|
|
104
111
|
T = t.TypeVar("T")
|
|
105
112
|
A = t.TypeVar("A", Scalar, NamedArray, jnp.ndarray)
|
|
106
113
|
|
|
@@ -337,6 +344,14 @@ def around(a: A) -> A:
|
|
|
337
344
|
return wrap_elemwise_unary(jnp.around, a)
|
|
338
345
|
|
|
339
346
|
|
|
347
|
+
def bitwise_count(a: A) -> A:
|
|
348
|
+
return wrap_elemwise_unary(jnp.bitwise_count, a)
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def bitwise_invert(a: A) -> A:
|
|
352
|
+
return wrap_elemwise_unary(jnp.bitwise_invert, a)
|
|
353
|
+
|
|
354
|
+
|
|
340
355
|
def bitwise_not(a: A) -> A:
|
|
341
356
|
return wrap_elemwise_unary(jnp.bitwise_not, a)
|
|
342
357
|
|
|
@@ -791,6 +806,7 @@ def argsort(a: NamedArray, axis: AxisSelector) -> NamedArray:
|
|
|
791
806
|
|
|
792
807
|
# elemwise binary ops
|
|
793
808
|
|
|
809
|
+
|
|
794
810
|
# Note that all the heavy lifting is done by the `wrap_elemwise_binary` decorator
|
|
795
811
|
@wrap_elemwise_binary
|
|
796
812
|
def add(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
@@ -816,6 +832,14 @@ def bitwise_and(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
|
816
832
|
return jnp.bitwise_and(x1, x2) # type: ignore
|
|
817
833
|
|
|
818
834
|
|
|
835
|
+
@wrap_elemwise_binary
|
|
836
|
+
def bitwise_left_shift(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
837
|
+
"""
|
|
838
|
+
Named version of [jax.numpy.bitwise_left_shift](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.bitwise_left_shift.html)
|
|
839
|
+
"""
|
|
840
|
+
return jnp.bitwise_left_shift(x1, x2) # type: ignore
|
|
841
|
+
|
|
842
|
+
|
|
819
843
|
@wrap_elemwise_binary
|
|
820
844
|
def bitwise_or(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
821
845
|
"""
|
|
@@ -824,6 +848,14 @@ def bitwise_or(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
|
824
848
|
return jnp.bitwise_or(x1, x2) # type: ignore
|
|
825
849
|
|
|
826
850
|
|
|
851
|
+
@wrap_elemwise_binary
|
|
852
|
+
def bitwise_right_shift(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
853
|
+
"""
|
|
854
|
+
Named version of [jax.numpy.bitwise_right_shift](https://jax.readthedocs.io/en/latest/_autosummary/jax.numpy.bitwise_right_shift.html)
|
|
855
|
+
"""
|
|
856
|
+
return jnp.bitwise_right_shift(x1, x2) # type: ignore
|
|
857
|
+
|
|
858
|
+
|
|
827
859
|
@wrap_elemwise_binary
|
|
828
860
|
def bitwise_xor(x1: NamedOrNumeric, x2: NamedOrNumeric, /) -> NamedOrNumeric:
|
|
829
861
|
"""
|
|
@@ -1080,6 +1112,8 @@ __all__ = [
|
|
|
1080
1112
|
"arctan",
|
|
1081
1113
|
"arctanh",
|
|
1082
1114
|
"around",
|
|
1115
|
+
"bitwise_count",
|
|
1116
|
+
"bitwise_invert",
|
|
1083
1117
|
"bitwise_not",
|
|
1084
1118
|
"cbrt",
|
|
1085
1119
|
"ceil",
|
|
@@ -1171,6 +1205,8 @@ __all__ = [
|
|
|
1171
1205
|
"unique_counts",
|
|
1172
1206
|
"unique_inverse",
|
|
1173
1207
|
"unique_all",
|
|
1208
|
+
"packbits",
|
|
1209
|
+
"unpackbits",
|
|
1174
1210
|
"searchsorted",
|
|
1175
1211
|
"bincount",
|
|
1176
1212
|
"clip",
|
|
@@ -1179,7 +1215,9 @@ __all__ = [
|
|
|
1179
1215
|
"add",
|
|
1180
1216
|
"arctan2",
|
|
1181
1217
|
"bitwise_and",
|
|
1218
|
+
"bitwise_left_shift",
|
|
1182
1219
|
"bitwise_or",
|
|
1220
|
+
"bitwise_right_shift",
|
|
1183
1221
|
"bitwise_xor",
|
|
1184
1222
|
"divide",
|
|
1185
1223
|
"divmod",
|
|
@@ -1,3 +1,8 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
# This whole file is copied from Equinox.
|
|
2
7
|
# (c) 2023, Google LLC. and/or Patrick Kidger. Apache 2.0 licensed.
|
|
3
8
|
# Patrick doesn't like that I depend on Equinox internals, so I copied this stuff
|
|
@@ -1,3 +1,8 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
import functools as ft
|
|
2
7
|
import typing
|
|
3
8
|
import warnings
|
|
@@ -30,8 +35,7 @@ def dot(
|
|
|
30
35
|
preferred_element_type: Optional[DTypeLike] = None,
|
|
31
36
|
out_axes: Optional[PartialAxisSpec] = ...,
|
|
32
37
|
dot_general=jax.lax.dot_general,
|
|
33
|
-
) -> NamedArray:
|
|
34
|
-
...
|
|
38
|
+
) -> NamedArray: ...
|
|
35
39
|
|
|
36
40
|
|
|
37
41
|
@typing.overload
|
|
@@ -42,8 +46,7 @@ def dot(
|
|
|
42
46
|
preferred_element_type: Optional[DTypeLike] = None,
|
|
43
47
|
out_axes: Optional[PartialAxisSpec] = ...,
|
|
44
48
|
dot_general=jax.lax.dot_general,
|
|
45
|
-
) -> NamedArray:
|
|
46
|
-
...
|
|
49
|
+
) -> NamedArray: ...
|
|
47
50
|
|
|
48
51
|
|
|
49
52
|
def dot(
|
|
@@ -1,9 +1,13 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
from functools import partial
|
|
2
7
|
|
|
3
8
|
from jax import custom_jvp, custom_vjp, lax
|
|
4
9
|
from jax import numpy as jnp
|
|
5
10
|
|
|
6
|
-
|
|
7
11
|
# All of this is copy paste from flax/linen/fp8_ops.py
|
|
8
12
|
# (Until we get to the module)
|
|
9
13
|
|
|
@@ -1,3 +1,8 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
import dataclasses
|
|
2
7
|
import functools as ft
|
|
3
8
|
import inspect
|
|
@@ -16,7 +21,6 @@ from haliax.core import NamedArray
|
|
|
16
21
|
from haliax.jax_utils import is_jax_array_like, multilevel_scan, tree_checkpoint_name
|
|
17
22
|
from haliax.util import is_jax_or_hax_array_like, is_named_array
|
|
18
23
|
|
|
19
|
-
|
|
20
24
|
BoolAxisSpec = Union[bool, Callable[[Any], bool]]
|
|
21
25
|
Carry = TypeVar("Carry")
|
|
22
26
|
X = TypeVar("X", contravariant=True)
|
|
@@ -31,8 +35,7 @@ def is_named_or_shaped_array_like(x):
|
|
|
31
35
|
class ScanFn(Protocol[Carry, Args, Y]):
|
|
32
36
|
""" """
|
|
33
37
|
|
|
34
|
-
def __call__(self, carry: Carry, *args: Args.args, **kwargs: Args.kwargs) -> tuple[Carry, Y]:
|
|
35
|
-
...
|
|
38
|
+
def __call__(self, carry: Carry, *args: Args.args, **kwargs: Args.kwargs) -> tuple[Carry, Y]: ...
|
|
36
39
|
|
|
37
40
|
|
|
38
41
|
@dataclasses.dataclass(frozen=True)
|
|
@@ -262,8 +265,7 @@ def scan(
|
|
|
262
265
|
reverse: bool = False,
|
|
263
266
|
unroll: int = 1,
|
|
264
267
|
is_scanned: BoolAxisSpec = is_named_or_shaped_array_like,
|
|
265
|
-
) -> Callable[[Carry, PyTree[X]], tuple[Carry, PyTree[Y]]]:
|
|
266
|
-
...
|
|
268
|
+
) -> Callable[[Carry, PyTree[X]], tuple[Carry, PyTree[Y]]]: ...
|
|
267
269
|
|
|
268
270
|
|
|
269
271
|
@overload
|
|
@@ -275,8 +277,7 @@ def scan(
|
|
|
275
277
|
reverse: bool = False,
|
|
276
278
|
unroll: int = 1,
|
|
277
279
|
is_scanned: BoolAxisSpec = is_named_or_shaped_array_like,
|
|
278
|
-
) -> Callable:
|
|
279
|
-
...
|
|
280
|
+
) -> Callable: ...
|
|
280
281
|
|
|
281
282
|
|
|
282
283
|
def scan(
|
|
@@ -444,8 +445,7 @@ def fold(
|
|
|
444
445
|
reverse: bool = False,
|
|
445
446
|
unroll: int = 1,
|
|
446
447
|
is_scanned: BoolAxisSpec = is_named_or_shaped_array_like,
|
|
447
|
-
) -> Callable[[Carry, PyTree[X]], Carry]:
|
|
448
|
-
...
|
|
448
|
+
) -> Callable[[Carry, PyTree[X]], Carry]: ...
|
|
449
449
|
|
|
450
450
|
|
|
451
451
|
@overload
|
|
@@ -457,8 +457,7 @@ def fold(
|
|
|
457
457
|
reverse: bool = False,
|
|
458
458
|
unroll: int = 1,
|
|
459
459
|
is_scanned: BoolAxisSpec = is_named_or_shaped_array_like,
|
|
460
|
-
) -> Callable:
|
|
461
|
-
...
|
|
460
|
+
) -> Callable: ...
|
|
462
461
|
|
|
463
462
|
|
|
464
463
|
def fold(
|
|
@@ -549,9 +548,7 @@ def _zero_if_array_else_none(x: Any) -> ResolvedUnnamedAxisSpec:
|
|
|
549
548
|
return 0 if is_jax_array_like(x) else None
|
|
550
549
|
|
|
551
550
|
|
|
552
|
-
def _format_tree_path(
|
|
553
|
-
path: tuple[jtu.KeyEntry, ...], arg_pos_names: dict[int, str] | None = None
|
|
554
|
-
) -> str:
|
|
551
|
+
def _format_tree_path(path: tuple[jtu.KeyEntry, ...], arg_pos_names: dict[int, str] | None = None) -> str:
|
|
555
552
|
parts: list[str] = []
|
|
556
553
|
i = 0
|
|
557
554
|
if len(path) >= 2 and isinstance(path[0], jtu.SequenceKey):
|
|
@@ -1,3 +1,8 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
# Module to support torch-style "state dict" serialization via safetensors
|
|
2
7
|
import dataclasses
|
|
3
8
|
import typing
|
|
@@ -19,7 +24,6 @@ from haliax.core import NamedArray, named
|
|
|
19
24
|
from haliax.jax_utils import is_jax_array_like, is_scalarish
|
|
20
25
|
from haliax.tree_util import scan_aware_tree_map
|
|
21
26
|
|
|
22
|
-
|
|
23
27
|
try:
|
|
24
28
|
import safetensors
|
|
25
29
|
except ImportError:
|
|
@@ -92,6 +96,7 @@ def _flatten_to_unflatten(t, state_dict, prefix):
|
|
|
92
96
|
"""
|
|
93
97
|
Flatten the torch compatible state_dict before loading into t, and then recover the unflattened layers.
|
|
94
98
|
"""
|
|
99
|
+
|
|
95
100
|
# typically, `t` is a bunch of ShapeDtypeStructs, which can't be transposed etc. so we instead have to zeros()
|
|
96
101
|
# into real arrays (that aren't actually real b/c this is inside a jit)
|
|
97
102
|
def _dt_struct_to_array(struct):
|
|
@@ -107,18 +112,15 @@ def _flatten_to_unflatten(t, state_dict, prefix):
|
|
|
107
112
|
|
|
108
113
|
|
|
109
114
|
@typing.overload
|
|
110
|
-
def with_prefix(prefix: str | None, leaf: str) -> str:
|
|
111
|
-
...
|
|
115
|
+
def with_prefix(prefix: str | None, leaf: str) -> str: ...
|
|
112
116
|
|
|
113
117
|
|
|
114
118
|
@typing.overload
|
|
115
|
-
def with_prefix(prefix: str, leaf: None) -> str:
|
|
116
|
-
...
|
|
119
|
+
def with_prefix(prefix: str, leaf: None) -> str: ...
|
|
117
120
|
|
|
118
121
|
|
|
119
122
|
@typing.overload
|
|
120
|
-
def with_prefix(prefix: Optional[str], leaf: Optional[str]) -> Optional[str]:
|
|
121
|
-
...
|
|
123
|
+
def with_prefix(prefix: Optional[str], leaf: Optional[str]) -> Optional[str]: ...
|
|
122
124
|
|
|
123
125
|
|
|
124
126
|
def with_prefix(prefix: Optional[str], leaf: Optional[str]) -> Optional[str]:
|
|
@@ -1,5 +1,9 @@
|
|
|
1
|
-
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
2
5
|
|
|
6
|
+
from typing import Callable, MutableMapping, Sequence, TypeAlias, TypeVar
|
|
3
7
|
|
|
4
8
|
T = TypeVar("T")
|
|
5
9
|
U = TypeVar("U")
|
|
@@ -1,3 +1,8 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
1
6
|
import typing
|
|
2
7
|
from dataclasses import dataclass
|
|
3
8
|
from math import prod
|
|
@@ -88,8 +93,7 @@ def selects_axis(selector: AxisSelection, selected: AxisSelection) -> bool:
|
|
|
88
93
|
return True
|
|
89
94
|
|
|
90
95
|
|
|
91
|
-
class _Sentinel:
|
|
92
|
-
...
|
|
96
|
+
class _Sentinel: ...
|
|
93
97
|
|
|
94
98
|
|
|
95
99
|
def is_axis_compatible(ax1: AxisSelector, ax2: AxisSelector):
|
|
@@ -140,23 +144,19 @@ def axis_spec_to_shape_dict(axis_spec: AxisSelection) -> dict[str, Optional[int]
|
|
|
140
144
|
|
|
141
145
|
|
|
142
146
|
@typing.overload
|
|
143
|
-
def axis_spec_to_tuple(axis_spec: ShapeDict) -> tuple[Axis, ...]:
|
|
144
|
-
...
|
|
147
|
+
def axis_spec_to_tuple(axis_spec: ShapeDict) -> tuple[Axis, ...]: ...
|
|
145
148
|
|
|
146
149
|
|
|
147
150
|
@typing.overload
|
|
148
|
-
def axis_spec_to_tuple(axis_spec: AxisSpec) -> tuple[Axis, ...]:
|
|
149
|
-
...
|
|
151
|
+
def axis_spec_to_tuple(axis_spec: AxisSpec) -> tuple[Axis, ...]: ...
|
|
150
152
|
|
|
151
153
|
|
|
152
154
|
@typing.overload
|
|
153
|
-
def axis_spec_to_tuple(axis_spec: PartialShapeDict) -> tuple[AxisSelector, ...]:
|
|
154
|
-
...
|
|
155
|
+
def axis_spec_to_tuple(axis_spec: PartialShapeDict) -> tuple[AxisSelector, ...]: ...
|
|
155
156
|
|
|
156
157
|
|
|
157
158
|
@typing.overload
|
|
158
|
-
def axis_spec_to_tuple(axis_spec: AxisSelection) -> tuple[AxisSelector, ...]:
|
|
159
|
-
...
|
|
159
|
+
def axis_spec_to_tuple(axis_spec: AxisSelection) -> tuple[AxisSelector, ...]: ...
|
|
160
160
|
|
|
161
161
|
|
|
162
162
|
def axis_spec_to_tuple(axis_spec: AxisSelection) -> tuple[AxisSelector, ...]:
|
|
@@ -230,23 +230,19 @@ def concat_axes(a1, a2):
|
|
|
230
230
|
|
|
231
231
|
|
|
232
232
|
@typing.overload
|
|
233
|
-
def union_axes(a1: ShapeDict, a2: AxisSpec) -> ShapeDict:
|
|
234
|
-
...
|
|
233
|
+
def union_axes(a1: ShapeDict, a2: AxisSpec) -> ShapeDict: ...
|
|
235
234
|
|
|
236
235
|
|
|
237
236
|
@typing.overload
|
|
238
|
-
def union_axes(a1: AxisSpec, a2: ShapeDict) -> ShapeDict:
|
|
239
|
-
...
|
|
237
|
+
def union_axes(a1: AxisSpec, a2: ShapeDict) -> ShapeDict: ...
|
|
240
238
|
|
|
241
239
|
|
|
242
240
|
@typing.overload
|
|
243
|
-
def union_axes(a1: AxisSpec, a2: AxisSpec) -> AxisSpec:
|
|
244
|
-
...
|
|
241
|
+
def union_axes(a1: AxisSpec, a2: AxisSpec) -> AxisSpec: ...
|
|
245
242
|
|
|
246
243
|
|
|
247
244
|
@typing.overload
|
|
248
|
-
def union_axes(a1: AxisSelection, a2: AxisSelection) -> AxisSelection:
|
|
249
|
-
...
|
|
245
|
+
def union_axes(a1: AxisSelection, a2: AxisSelection) -> AxisSelection: ...
|
|
250
246
|
|
|
251
247
|
|
|
252
248
|
def union_axes(a1: AxisSelection, a2: AxisSelection) -> AxisSelection:
|
|
@@ -372,23 +368,19 @@ def without_axes(axis_spec: AxisSelection, to_remove: AxisSelection, allow_misma
|
|
|
372
368
|
|
|
373
369
|
|
|
374
370
|
@typing.overload
|
|
375
|
-
def unsize_axes(axis_spec: PartialShapeDict, to_unsize: AxisSelection) -> PartialShapeDict:
|
|
376
|
-
...
|
|
371
|
+
def unsize_axes(axis_spec: PartialShapeDict, to_unsize: AxisSelection) -> PartialShapeDict: ...
|
|
377
372
|
|
|
378
373
|
|
|
379
374
|
@typing.overload
|
|
380
|
-
def unsize_axes(axis_spec: AxisSelection, to_unsize: AxisSelection) -> AxisSelection:
|
|
381
|
-
...
|
|
375
|
+
def unsize_axes(axis_spec: AxisSelection, to_unsize: AxisSelection) -> AxisSelection: ...
|
|
382
376
|
|
|
383
377
|
|
|
384
378
|
@typing.overload
|
|
385
|
-
def unsize_axes(axis_spec: PartialShapeDict) -> PartialShapeDict:
|
|
386
|
-
...
|
|
379
|
+
def unsize_axes(axis_spec: PartialShapeDict) -> PartialShapeDict: ...
|
|
387
380
|
|
|
388
381
|
|
|
389
382
|
@typing.overload
|
|
390
|
-
def unsize_axes(axis_spec: AxisSelection) -> AxisSelection:
|
|
391
|
-
...
|
|
383
|
+
def unsize_axes(axis_spec: AxisSelection) -> AxisSelection: ...
|
|
392
384
|
|
|
393
385
|
|
|
394
386
|
def unsize_axes(axis_spec: AxisSelection, to_unsize: Optional[AxisSelection] = None) -> AxisSelection:
|
|
@@ -424,13 +416,11 @@ def unsize_axes(axis_spec: AxisSelection, to_unsize: Optional[AxisSelection] = N
|
|
|
424
416
|
|
|
425
417
|
|
|
426
418
|
@overload
|
|
427
|
-
def replace_axis(axis_spec: AxisSpec, old: AxisSelector, new: AxisSpec) -> AxisSpec:
|
|
428
|
-
...
|
|
419
|
+
def replace_axis(axis_spec: AxisSpec, old: AxisSelector, new: AxisSpec) -> AxisSpec: ...
|
|
429
420
|
|
|
430
421
|
|
|
431
422
|
@overload
|
|
432
|
-
def replace_axis(axis_spec: AxisSelection, old: AxisSelector, new: AxisSelection) -> AxisSelection:
|
|
433
|
-
...
|
|
423
|
+
def replace_axis(axis_spec: AxisSelection, old: AxisSelector, new: AxisSelection) -> AxisSelection: ...
|
|
434
424
|
|
|
435
425
|
|
|
436
426
|
def replace_axis(axis_spec: AxisSelection, old: AxisSelector, new: AxisSelection) -> AxisSelection:
|
|
@@ -466,13 +456,11 @@ def replace_axis(axis_spec: AxisSelection, old: AxisSelector, new: AxisSelection
|
|
|
466
456
|
|
|
467
457
|
|
|
468
458
|
@overload
|
|
469
|
-
def intersect_axes(ax1: ShapeDict, ax2: AxisSelection) -> ShapeDict:
|
|
470
|
-
...
|
|
459
|
+
def intersect_axes(ax1: ShapeDict, ax2: AxisSelection) -> ShapeDict: ...
|
|
471
460
|
|
|
472
461
|
|
|
473
462
|
@overload
|
|
474
|
-
def intersect_axes(ax1: tuple[AxisSelector, ...], ax2: AxisSpec) -> tuple[Axis, ...]:
|
|
475
|
-
...
|
|
463
|
+
def intersect_axes(ax1: tuple[AxisSelector, ...], ax2: AxisSpec) -> tuple[Axis, ...]: ...
|
|
476
464
|
|
|
477
465
|
|
|
478
466
|
@overload
|
|
@@ -550,13 +538,11 @@ def axis_size(ax: AxisSpec) -> int:
|
|
|
550
538
|
|
|
551
539
|
|
|
552
540
|
@typing.overload
|
|
553
|
-
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelector) -> Axis:
|
|
554
|
-
...
|
|
541
|
+
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelector) -> Axis: ...
|
|
555
542
|
|
|
556
543
|
|
|
557
544
|
@typing.overload
|
|
558
|
-
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
|
|
559
|
-
...
|
|
545
|
+
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec: ...
|
|
560
546
|
|
|
561
547
|
|
|
562
548
|
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
|