haliax 1.4.dev398__tar.gz → 1.4.dev400__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.dev398 → haliax-1.4.dev400}/PKG-INFO +1 -1
  2. haliax-1.4.dev400/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/core.py +15 -7
  4. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/jax_utils.py +15 -0
  5. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/ops.py +2 -2
  6. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/wrap.py +3 -3
  7. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_ops.py +25 -0
  8. haliax-1.4.dev398/src/haliax/__about__.py +0 -1
  9. {haliax-1.4.dev398 → haliax-1.4.dev400}/.coveragerc +0 -0
  10. {haliax-1.4.dev398 → haliax-1.4.dev400}/.flake8 +0 -0
  11. {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/publish_dev.yaml +0 -0
  12. {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_pre_commit.yaml +0 -0
  13. {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  14. {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_tests.yaml +0 -0
  15. {haliax-1.4.dev398 → haliax-1.4.dev400}/.gitignore +0 -0
  16. {haliax-1.4.dev398 → haliax-1.4.dev400}/.playbooks/add-types.md +0 -0
  17. {haliax-1.4.dev398 → haliax-1.4.dev400}/.playbooks/wrap-non-named.md +0 -0
  18. {haliax-1.4.dev398 → haliax-1.4.dev400}/.pre-commit-config.yaml +0 -0
  19. {haliax-1.4.dev398 → haliax-1.4.dev400}/.readthedocs.yaml +0 -0
  20. {haliax-1.4.dev398 → haliax-1.4.dev400}/AGENTS.md +0 -0
  21. {haliax-1.4.dev398 → haliax-1.4.dev400}/CONTRIBUTING.md +0 -0
  22. {haliax-1.4.dev398 → haliax-1.4.dev400}/LICENSE +0 -0
  23. {haliax-1.4.dev398 → haliax-1.4.dev400}/README.md +0 -0
  24. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/api.md +0 -0
  25. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/broadcasting.md +0 -0
  26. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/cheatsheet.md +0 -0
  27. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/css/material.css +0 -0
  28. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/css/mkdocstrings.css +0 -0
  29. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/faq.md +0 -0
  30. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/data_parallel_mesh.png +0 -0
  31. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  32. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_1d.png +0 -0
  33. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_1d_zero.png +0 -0
  34. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d.png +0 -0
  35. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  36. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  37. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  38. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  39. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_zero.png +0 -0
  40. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/fp8.md +0 -0
  41. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/index.md +0 -0
  42. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/indexing.md +0 -0
  43. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/matmul.md +0 -0
  44. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/nn.md +0 -0
  45. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/partitioning.md +0 -0
  46. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/primer.md +0 -0
  47. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/rearrange.ipynb +0 -0
  48. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/rearrange.md +0 -0
  49. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/requirements.txt +0 -0
  50. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/scan.md +0 -0
  51. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/state-dict.md +0 -0
  52. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/tutorial.md +0 -0
  53. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/typing.md +0 -0
  54. {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/vmap.md +0 -0
  55. {haliax-1.4.dev398 → haliax-1.4.dev400}/mkdocs.yml +0 -0
  56. {haliax-1.4.dev398 → haliax-1.4.dev400}/pyproject.toml +0 -0
  57. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/__init__.py +0 -0
  58. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/__init__.py +0 -0
  59. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/compile_utils.py +0 -0
  60. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/dot.py +0 -0
  61. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/einsum.py +0 -0
  62. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/fp8.py +0 -0
  63. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/parsing.py +0 -0
  64. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/rearrange.py +0 -0
  65. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/scan.py +0 -0
  66. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/state_dict.py +0 -0
  67. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/util.py +0 -0
  68. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/axis.py +0 -0
  69. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/debug.py +0 -0
  70. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/haxtyping.py +0 -0
  71. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/hof.py +0 -0
  72. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/__init__.py +0 -0
  73. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/activations.py +0 -0
  74. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/attention.py +0 -0
  75. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/conv.py +0 -0
  76. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/dropout.py +0 -0
  77. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/embedding.py +0 -0
  78. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/linear.py +0 -0
  79. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/loss.py +0 -0
  80. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/mlp.py +0 -0
  81. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/normalization.py +0 -0
  82. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/pool.py +0 -0
  83. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/scan.py +0 -0
  84. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/partitioning.py +0 -0
  85. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/quantization.py +0 -0
  86. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/random.py +0 -0
  87. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/specialized_fns.py +0 -0
  88. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/state_dict.py +0 -0
  89. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/tree_util.py +0 -0
  90. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/types.py +0 -0
  91. {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/util.py +0 -0
  92. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/core_test.py +0 -0
  93. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_attention.py +0 -0
  94. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_axis.py +0 -0
  95. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_conv.py +0 -0
  96. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_debug.py +0 -0
  97. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_dot.py +0 -0
  98. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_dtype_typing.py +0 -0
  99. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_einsum.py +0 -0
  100. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_fp8.py +0 -0
  101. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_hof.py +0 -0
  102. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_int8.py +0 -0
  103. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_moe_linear.py +0 -0
  104. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_namedarray_typing.py +0 -0
  105. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_nn.py +0 -0
  106. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_parsing.py +0 -0
  107. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_partitioning.py +0 -0
  108. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_pool.py +0 -0
  109. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_random.py +0 -0
  110. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_rearrange.py +0 -0
  111. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_scan.py +0 -0
  112. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_scatter_gather.py +0 -0
  113. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_specialized_fns.py +0 -0
  114. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_state_dict.py +0 -0
  115. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_tree_util.py +0 -0
  116. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_utils.py +0 -0
  117. {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_visualize_sharding.py +0 -0
  118. {haliax-1.4.dev398 → haliax-1.4.dev400}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev398
3
+ Version: 1.4.dev400
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.dev400"
@@ -15,7 +15,7 @@ import numpy as np
15
15
 
16
16
  import haliax
17
17
  import haliax.axis
18
- from haliax.jax_utils import is_jax_array_like, is_pallas_dslice
18
+ from haliax.jax_utils import ensure_scalar, is_jax_array_like, is_pallas_dslice
19
19
  from haliax.util import ensure_tuple
20
20
 
21
21
  from ._src.util import index_where, py_slice, slice_t
@@ -1115,9 +1115,7 @@ def updated_slice(
1115
1115
  if axis_index is None:
1116
1116
  raise ValueError(f"axis {axis} not found in {array}")
1117
1117
  if isinstance(s, NamedArray): # this can happen in the vmap case
1118
- if s.ndim != 0:
1119
- raise ValueError(f"NamedArray {s} must be a scalar for axis {axis} in updated_slice")
1120
- s = s.scalar()
1118
+ s = ensure_scalar(s, name=str(axis))
1121
1119
 
1122
1120
  array_slice_indices[axis_index] = s
1123
1121
  total_length = array.axes[axis_index].size
@@ -1361,13 +1359,23 @@ def unbind(array: NamedArray, axis: AxisSelector) -> List[NamedArray]:
1361
1359
  return [haliax.auto_sharded(NamedArray(a, new_axes)) for a in arrays]
1362
1360
 
1363
1361
 
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
1362
+ def roll(
1363
+ array: NamedArray,
1364
+ shift: Union[IntScalar, Tuple[int, ...], "NamedArray"],
1365
+ axis: AxisSelection,
1366
+ ) -> NamedArray:
1367
+ """Roll an array along an axis or axes.
1368
+
1369
+ ``shift`` may be a scalar ``NamedArray`` in addition to an ``int`` or tuple of
1370
+ integers.
1367
1371
  """
1372
+
1368
1373
  axis_indices = array.axis_indices(axis)
1369
1374
  if axis_indices is None:
1370
1375
  raise ValueError(f"axis {axis} not found in {array}")
1376
+
1377
+ shift = ensure_scalar(shift, name="shift")
1378
+
1371
1379
  return NamedArray(jnp.roll(array.array, shift, axis_indices), array.axes)
1372
1380
 
1373
1381
 
@@ -153,6 +153,21 @@ def is_scalarish(x):
153
153
  return jnp.isscalar(x) or (getattr(x, "shape", None) == ())
154
154
 
155
155
 
156
+ def ensure_scalar(x, *, name: str = "value"):
157
+ """Return ``x`` if it is not a :class:`NamedArray`, otherwise ensure it is a scalar.
158
+
159
+ This is useful for APIs that can accept either Python scalars or scalar
160
+ ``NamedArray`` objects (for example ``roll`` or ``updated_slice``). If ``x``
161
+ is a ``NamedArray`` with rank greater than 0 a :class:`TypeError` is raised.
162
+ """
163
+
164
+ if isinstance(x, haliax.NamedArray):
165
+ if x.ndim != 0:
166
+ raise TypeError(f"{name} must be a scalar NamedArray")
167
+ return x.array
168
+ return x
169
+
170
+
156
171
  def is_on_mac_metal():
157
172
  return jax.devices()[0].platform.lower() == "metal"
158
173
 
@@ -9,7 +9,7 @@ import haliax
9
9
 
10
10
  from .axis import Axis, AxisSelector, axis_name
11
11
  from .core import NamedArray, NamedOrNumeric, broadcast_arrays, broadcast_arrays_and_return_axes, named
12
- from .jax_utils import is_scalarish
12
+ from .jax_utils import ensure_scalar, is_scalarish
13
13
 
14
14
 
15
15
  def trace(array: NamedArray, axis1: AxisSelector, axis2: AxisSelector, offset=0, dtype=None) -> NamedArray:
@@ -89,7 +89,7 @@ def where(
89
89
  x = named(x, ())
90
90
  x, y = broadcast_arrays(x, y)
91
91
  if isinstance(condition, NamedArray):
92
- condition = condition.scalar()
92
+ condition = ensure_scalar(condition, name="condition")
93
93
  return jax.lax.cond(condition, lambda _: x, lambda _: y, None)
94
94
 
95
95
  condition, x, y = broadcast_arrays(condition, x, y) # type: ignore
@@ -5,7 +5,7 @@ import jax
5
5
  from haliax.core import NamedArray, _broadcast_order, broadcast_to
6
6
 
7
7
  from .axis import AxisSelection, AxisSelector, axis_spec_to_shape_dict, eliminate_axes
8
- from .jax_utils import is_scalarish
8
+ from .jax_utils import ensure_scalar, is_scalarish
9
9
 
10
10
 
11
11
  def wrap_elemwise_unary(f, a, *args, **kwargs):
@@ -105,7 +105,7 @@ def wrap_elemwise_binary(op):
105
105
  else:
106
106
  if is_scalarish(b):
107
107
  return NamedArray(op(a.array, b), a.axes)
108
- a = a.scalar()
108
+ a = ensure_scalar(a)
109
109
  return op(a, b)
110
110
 
111
111
  return NamedArray(op(a.array, b), a.axes)
@@ -119,7 +119,7 @@ def wrap_elemwise_binary(op):
119
119
  else:
120
120
  if is_scalarish(a):
121
121
  return NamedArray(op(a, b.array), b.axes)
122
- b = b.scalar()
122
+ b = ensure_scalar(b)
123
123
  return op(a, b)
124
124
 
125
125
  return NamedArray(op(a, b.array), b.axes)
@@ -423,3 +423,28 @@ 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)
440
+
441
+
442
+ def test_roll_bad_named_shift():
443
+ H = Axis("H", 4)
444
+ W = Axis("W", 3)
445
+
446
+ arr = hax.arange((H, W))
447
+ shift = hax.arange((Axis("dummy", 2),))
448
+
449
+ with pytest.raises(TypeError):
450
+ hax.roll(arr, shift, H)
@@ -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