haliax 1.4.dev397__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.dev397 → haliax-1.4.dev399}/PKG-INFO +1 -1
- haliax-1.4.dev399/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/core.py +16 -3
- haliax-1.4.dev399/tests/test_moe_linear.py +70 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_ops.py +14 -0
- haliax-1.4.dev397/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.coveragerc +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.flake8 +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.gitignore +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/AGENTS.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/LICENSE +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/README.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/api.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/css/material.css +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/faq.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/fp8.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/index.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/indexing.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/matmul.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/nn.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/partitioning.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/primer.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/rearrange.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/requirements.txt +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/scan.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/state-dict.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/tutorial.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/typing.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/vmap.md +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/mkdocs.yml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/pyproject.toml +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/random.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/types.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/util.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/core_test.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_attention.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_axis.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_conv.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_debug.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_dot.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_hof.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_int8.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_nn.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_pool.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_random.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_scan.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_utils.py +0 -0
- {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev397 → 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
|
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
import jax
|
|
2
|
+
import jax.random as jrandom
|
|
3
|
+
from jax import numpy as jnp
|
|
4
|
+
|
|
5
|
+
import haliax as hax
|
|
6
|
+
from haliax.nn import MoELinear
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def _expected_moe_linear_output(moe: MoELinear, x: hax.NamedArray, group_sizes: hax.NamedArray):
|
|
10
|
+
dim_numbers = jax.lax.RaggedDotDimensionNumbers(
|
|
11
|
+
(
|
|
12
|
+
((x.axis_indices(moe.In),), (moe.weight.axis_indices(moe.In),)),
|
|
13
|
+
((), ()),
|
|
14
|
+
),
|
|
15
|
+
x.axis_indices(hax.axis.without_axes(x.axes, moe.In)),
|
|
16
|
+
(moe.weight.axis_indices(moe.Experts),),
|
|
17
|
+
)
|
|
18
|
+
out_raw = jax.lax.ragged_dot_general(
|
|
19
|
+
lhs=x.array,
|
|
20
|
+
rhs=moe.weight.array,
|
|
21
|
+
group_sizes=group_sizes.array,
|
|
22
|
+
ragged_dot_dimension_numbers=dim_numbers,
|
|
23
|
+
)
|
|
24
|
+
out_axes = hax.replace_axis(x.axes, moe.In, moe.Out)
|
|
25
|
+
out = hax.named(out_raw, out_axes)
|
|
26
|
+
if moe.bias is not None:
|
|
27
|
+
out = out + moe.bias
|
|
28
|
+
return out
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def test_moe_linear_matches_ragged_dot_general():
|
|
32
|
+
B, In, Out, E = hax.make_axes(B=3, In=4, Out=5, E=2)
|
|
33
|
+
key = jrandom.PRNGKey(0)
|
|
34
|
+
moe = MoELinear.init(E, In, Out, key=key)
|
|
35
|
+
|
|
36
|
+
x = hax.random.normal(jrandom.PRNGKey(1), (B, In))
|
|
37
|
+
group_sizes = hax.named(jnp.array([2, 1], dtype=jnp.int32), (E,))
|
|
38
|
+
|
|
39
|
+
actual = moe(x, group_sizes)
|
|
40
|
+
expected = _expected_moe_linear_output(moe, x, group_sizes)
|
|
41
|
+
|
|
42
|
+
assert actual.axes == expected.axes
|
|
43
|
+
assert jnp.allclose(actual.array, expected.array, rtol=1e-5, atol=1e-5)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def test_moe_linear_out_first_property():
|
|
47
|
+
E, In, Out = hax.make_axes(E=2, In=4, Out=3)
|
|
48
|
+
moe = MoELinear.init(E, In, Out, key=jrandom.PRNGKey(0), out_first=True)
|
|
49
|
+
assert moe.out_first
|
|
50
|
+
assert moe.weight.axes[:3] == (E, Out, In)
|
|
51
|
+
|
|
52
|
+
moe2 = MoELinear.init(E, In, Out, key=jrandom.PRNGKey(1), out_first=False)
|
|
53
|
+
assert not moe2.out_first
|
|
54
|
+
assert moe2.weight.axes[:3] == (E, In, Out)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def test_moe_linear_gmm_matches_ragged_dot_general():
|
|
58
|
+
B, In, Out, E = hax.make_axes(B=3, In=4, Out=5, E=2)
|
|
59
|
+
moe = MoELinear.init(E, In, Out, key=jrandom.PRNGKey(0), use_gmm=True)
|
|
60
|
+
|
|
61
|
+
x = hax.random.normal(jrandom.PRNGKey(1), (B, In))
|
|
62
|
+
group_sizes = hax.named(jnp.array([2, 1], dtype=jnp.int32), (E,))
|
|
63
|
+
|
|
64
|
+
with jax.sharding.Mesh(jax.devices(), ("data",)):
|
|
65
|
+
actual = moe(x, group_sizes)
|
|
66
|
+
|
|
67
|
+
expected = _expected_moe_linear_output(moe, x, group_sizes)
|
|
68
|
+
|
|
69
|
+
assert actual.axes == expected.axes
|
|
70
|
+
assert jnp.allclose(actual.array, expected.array, rtol=1e-5, atol=1e-5)
|
|
@@ -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.dev397"
|
|
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
|