haliax 1.4.dev313__tar.gz → 1.4.dev314__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.dev313 → haliax-1.4.dev314}/PKG-INFO +1 -1
- haliax-1.4.dev314/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/core.py +1 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/core_test.py +17 -0
- haliax-1.4.dev313/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.coveragerc +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.flake8 +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.gitignore +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/LICENSE +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/README.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/api.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/css/material.css +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/faq.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/fp8.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/hof.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/index.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/indexing.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/matmul.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/nn.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/partitioning.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/rearrange.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/requirements.txt +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/docs/tutorial.md +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/mkdocs.yml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/pyproject.toml +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/random.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/types.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/util.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_attention.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_axis.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_conv.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_debug.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_dot.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_hof.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_nn.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_ops.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_pool.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_random.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_scan.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev313 → haliax-1.4.dev314}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev314
|
|
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/
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev314"
|
|
@@ -952,6 +952,7 @@ def _compute_new_axes_and_slices_for_index(
|
|
|
952
952
|
# we allow this if it's a 0-d or 1-d array
|
|
953
953
|
if slice_.ndim == 0:
|
|
954
954
|
ordered_slices[axis_index] = slice_
|
|
955
|
+
kept_axes[axis_index] = False
|
|
955
956
|
elif slice_.ndim == 1:
|
|
956
957
|
# we allow this if it's a 1-d array, in which case we treat it as sugar for NamedArray(slice_, sliced-axis)
|
|
957
958
|
ordered_slices[axis_index] = haliax.named(slice_, axis_name(axis))
|
|
@@ -543,6 +543,23 @@ def test_index():
|
|
|
543
543
|
assert jnp.all(jnp.equal(named1[{"H": 0, "W": 0, "D": 0}], named1.array[0, 0, 0]))
|
|
544
544
|
|
|
545
545
|
|
|
546
|
+
def test_index_with_tracer():
|
|
547
|
+
H = Axis("H", 20)
|
|
548
|
+
W = Axis("W", 30)
|
|
549
|
+
D = Axis("D", 40)
|
|
550
|
+
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
551
|
+
|
|
552
|
+
@jax.jit
|
|
553
|
+
def f(idx):
|
|
554
|
+
return named1["H", idx]
|
|
555
|
+
|
|
556
|
+
idx = jnp.array([1, 2, 3])
|
|
557
|
+
assert jnp.all(jnp.equal(f(idx).array, named1.array[1:4, :, :]))
|
|
558
|
+
|
|
559
|
+
idx = jnp.array(0)
|
|
560
|
+
assert jnp.all(jnp.equal(f(idx).array, named1.array[0, :, :]))
|
|
561
|
+
|
|
562
|
+
|
|
546
563
|
def test_index_array_slices():
|
|
547
564
|
# fancier tests with array slices with named array args
|
|
548
565
|
H = Axis("H", 10)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev313"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|