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