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.
- {haliax-1.4.dev398 → haliax-1.4.dev400}/PKG-INFO +1 -1
- haliax-1.4.dev400/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/core.py +15 -7
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/jax_utils.py +15 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/ops.py +2 -2
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/wrap.py +3 -3
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_ops.py +25 -0
- haliax-1.4.dev398/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.coveragerc +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.flake8 +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.gitignore +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/AGENTS.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/LICENSE +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/README.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/api.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/css/material.css +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/faq.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/fp8.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/index.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/indexing.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/matmul.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/nn.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/partitioning.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/primer.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/rearrange.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/requirements.txt +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/scan.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/state-dict.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/tutorial.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/typing.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/docs/vmap.md +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/mkdocs.yml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/pyproject.toml +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/random.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/types.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/src/haliax/util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/core_test.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_attention.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_axis.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_conv.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_debug.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_dot.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_hof.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_int8.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_nn.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_pool.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_random.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_scan.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_utils.py +0 -0
- {haliax-1.4.dev398 → haliax-1.4.dev400}/tests/test_visualize_sharding.py +0 -0
- {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.
|
|
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
|
-
|
|
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(
|
|
1365
|
-
|
|
1366
|
-
|
|
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
|
|
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
|
|
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
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|