haliax 1.4.dev398__tar.gz → 1.4.dev399__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.dev398 → haliax-1.4.dev399}/PKG-INFO +1 -1
- haliax-1.4.dev399/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/core.py +16 -3
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_ops.py +14 -0
- haliax-1.4.dev398/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.coveragerc +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.flake8 +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.gitignore +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/AGENTS.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/LICENSE +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/README.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/api.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/css/material.css +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/faq.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/fp8.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/index.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/indexing.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/matmul.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/nn.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/partitioning.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/primer.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/rearrange.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/requirements.txt +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/scan.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/state-dict.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/tutorial.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/typing.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/docs/vmap.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/mkdocs.yml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/pyproject.toml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/random.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/types.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/core_test.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_attention.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_axis.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_conv.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_debug.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_dot.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_hof.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_int8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_nn.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_pool.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_random.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_utils.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev399}/uv.lock +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev399
|
|
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.dev399"
|
|
@@ -1361,13 +1361,26 @@ def unbind(array: NamedArray, axis: AxisSelector) -> List[NamedArray]:
|
|
|
1361
1361
|
return [haliax.auto_sharded(NamedArray(a, new_axes)) for a in arrays]
|
|
1362
1362
|
|
|
1363
1363
|
|
|
1364
|
-
def roll(
|
|
1365
|
-
|
|
1366
|
-
|
|
1364
|
+
def roll(
|
|
1365
|
+
array: NamedArray,
|
|
1366
|
+
shift: Union[IntScalar, Tuple[int, ...], "NamedArray"],
|
|
1367
|
+
axis: AxisSelection,
|
|
1368
|
+
) -> NamedArray:
|
|
1369
|
+
"""Roll an array along an axis or axes.
|
|
1370
|
+
|
|
1371
|
+
``shift`` may be a scalar ``NamedArray`` in addition to an ``int`` or tuple of
|
|
1372
|
+
integers.
|
|
1367
1373
|
"""
|
|
1374
|
+
|
|
1368
1375
|
axis_indices = array.axis_indices(axis)
|
|
1369
1376
|
if axis_indices is None:
|
|
1370
1377
|
raise ValueError(f"axis {axis} not found in {array}")
|
|
1378
|
+
|
|
1379
|
+
if isinstance(shift, NamedArray):
|
|
1380
|
+
if shift.ndim != 0:
|
|
1381
|
+
raise TypeError("shift must be a scalar NamedArray")
|
|
1382
|
+
shift = shift.array
|
|
1383
|
+
|
|
1371
1384
|
return NamedArray(jnp.roll(array.array, shift, axis_indices), array.axes)
|
|
1372
1385
|
|
|
1373
1386
|
|
|
@@ -423,3 +423,17 @@ def test_bincount():
|
|
|
423
423
|
out_w = hax.bincount(x, B, weights=w)
|
|
424
424
|
expected_w = jnp.bincount(x.array, weights=w.array, length=B.size)
|
|
425
425
|
assert jnp.allclose(out_w.array, expected_w)
|
|
426
|
+
|
|
427
|
+
|
|
428
|
+
def test_roll_scalar_named_shift():
|
|
429
|
+
H = Axis("H", 4)
|
|
430
|
+
W = Axis("W", 3)
|
|
431
|
+
|
|
432
|
+
arr = hax.arange((H, W))
|
|
433
|
+
shift = hax.named(jnp.array(1), ())
|
|
434
|
+
|
|
435
|
+
rolled = hax.roll(arr, shift, H)
|
|
436
|
+
expected = jnp.roll(arr.array, shift.array, axis=0)
|
|
437
|
+
|
|
438
|
+
assert rolled.axes == arr.axes
|
|
439
|
+
assert jnp.all(rolled.array == expected)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev398"
|
|
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
|
|
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
|