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.
Files changed (118) hide show
  1. {haliax-1.4.dev397 → haliax-1.4.dev399}/PKG-INFO +1 -1
  2. haliax-1.4.dev399/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/core.py +16 -3
  4. haliax-1.4.dev399/tests/test_moe_linear.py +70 -0
  5. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_ops.py +14 -0
  6. haliax-1.4.dev397/src/haliax/__about__.py +0 -1
  7. {haliax-1.4.dev397 → haliax-1.4.dev399}/.coveragerc +0 -0
  8. {haliax-1.4.dev397 → haliax-1.4.dev399}/.flake8 +0 -0
  9. {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/publish_dev.yaml +0 -0
  10. {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_pre_commit.yaml +0 -0
  11. {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  12. {haliax-1.4.dev397 → haliax-1.4.dev399}/.github/workflows/run_tests.yaml +0 -0
  13. {haliax-1.4.dev397 → haliax-1.4.dev399}/.gitignore +0 -0
  14. {haliax-1.4.dev397 → haliax-1.4.dev399}/.playbooks/add-types.md +0 -0
  15. {haliax-1.4.dev397 → haliax-1.4.dev399}/.playbooks/wrap-non-named.md +0 -0
  16. {haliax-1.4.dev397 → haliax-1.4.dev399}/.pre-commit-config.yaml +0 -0
  17. {haliax-1.4.dev397 → haliax-1.4.dev399}/.readthedocs.yaml +0 -0
  18. {haliax-1.4.dev397 → haliax-1.4.dev399}/AGENTS.md +0 -0
  19. {haliax-1.4.dev397 → haliax-1.4.dev399}/CONTRIBUTING.md +0 -0
  20. {haliax-1.4.dev397 → haliax-1.4.dev399}/LICENSE +0 -0
  21. {haliax-1.4.dev397 → haliax-1.4.dev399}/README.md +0 -0
  22. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/api.md +0 -0
  23. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/faq.md +0 -0
  28. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/index.md +0 -0
  40. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/nn.md +0 -0
  43. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/primer.md +0 -0
  45. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/rearrange.ipynb +0 -0
  46. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/rearrange.md +0 -0
  47. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/requirements.txt +0 -0
  48. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/scan.md +0 -0
  49. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/state-dict.md +0 -0
  50. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/tutorial.md +0 -0
  51. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/typing.md +0 -0
  52. {haliax-1.4.dev397 → haliax-1.4.dev399}/docs/vmap.md +0 -0
  53. {haliax-1.4.dev397 → haliax-1.4.dev399}/mkdocs.yml +0 -0
  54. {haliax-1.4.dev397 → haliax-1.4.dev399}/pyproject.toml +0 -0
  55. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/__init__.py +0 -0
  56. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/__init__.py +0 -0
  57. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/compile_utils.py +0 -0
  58. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/dot.py +0 -0
  59. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/einsum.py +0 -0
  60. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/fp8.py +0 -0
  61. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/parsing.py +0 -0
  62. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/rearrange.py +0 -0
  63. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/scan.py +0 -0
  64. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/state_dict.py +0 -0
  65. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/_src/util.py +0 -0
  66. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/axis.py +0 -0
  67. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/debug.py +0 -0
  68. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/haxtyping.py +0 -0
  69. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/hof.py +0 -0
  70. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/jax_utils.py +0 -0
  71. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/__init__.py +0 -0
  72. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/activations.py +0 -0
  73. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/attention.py +0 -0
  74. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/conv.py +0 -0
  75. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/dropout.py +0 -0
  76. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/embedding.py +0 -0
  77. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/linear.py +0 -0
  78. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/loss.py +0 -0
  79. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/mlp.py +0 -0
  80. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/normalization.py +0 -0
  81. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/pool.py +0 -0
  82. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/nn/scan.py +0 -0
  83. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/ops.py +0 -0
  84. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/partitioning.py +0 -0
  85. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/quantization.py +0 -0
  86. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/random.py +0 -0
  87. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/specialized_fns.py +0 -0
  88. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/state_dict.py +0 -0
  89. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/tree_util.py +0 -0
  90. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/types.py +0 -0
  91. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/util.py +0 -0
  92. {haliax-1.4.dev397 → haliax-1.4.dev399}/src/haliax/wrap.py +0 -0
  93. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/core_test.py +0 -0
  94. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_attention.py +0 -0
  95. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_axis.py +0 -0
  96. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_conv.py +0 -0
  97. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_debug.py +0 -0
  98. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_dot.py +0 -0
  99. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_dtype_typing.py +0 -0
  100. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_einsum.py +0 -0
  101. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_fp8.py +0 -0
  102. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_hof.py +0 -0
  103. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_int8.py +0 -0
  104. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_namedarray_typing.py +0 -0
  105. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_nn.py +0 -0
  106. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_parsing.py +0 -0
  107. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_partitioning.py +0 -0
  108. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_pool.py +0 -0
  109. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_random.py +0 -0
  110. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_rearrange.py +0 -0
  111. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_scan.py +0 -0
  112. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_scatter_gather.py +0 -0
  113. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_specialized_fns.py +0 -0
  114. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_state_dict.py +0 -0
  115. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_tree_util.py +0 -0
  116. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_utils.py +0 -0
  117. {haliax-1.4.dev397 → haliax-1.4.dev399}/tests/test_visualize_sharding.py +0 -0
  118. {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.dev397
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(array: NamedArray, shift: Union[int, Tuple[int, ...]], axis: AxisSelection) -> NamedArray:
1365
- """
1366
- Roll an array along an axis or axes. Analogous to np.roll
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