haliax 1.4.dev392__tar.gz → 1.4.dev393__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.dev392 → haliax-1.4.dev393}/PKG-INFO +1 -1
- haliax-1.4.dev393/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/core.py +41 -17
- haliax-1.4.dev392/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.coveragerc +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.flake8 +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.gitignore +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/AGENTS.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/LICENSE +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/README.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/api.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/css/material.css +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/faq.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/fp8.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/index.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/indexing.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/matmul.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/nn.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/partitioning.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/rearrange.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/requirements.txt +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/scan.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/state-dict.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/tutorial.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/typing.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/docs/vmap.md +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/mkdocs.yml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/pyproject.toml +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/random.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/types.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/util.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/core_test.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_attention.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_axis.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_conv.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_debug.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_dot.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_hof.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_int8.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_nn.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_ops.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_pool.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_random.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_scan.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_utils.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev392 → haliax-1.4.dev393}/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.dev393
|
|
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.dev393"
|
|
@@ -32,6 +32,7 @@ from .axis import (
|
|
|
32
32
|
dslice,
|
|
33
33
|
eliminate_axes,
|
|
34
34
|
selects_axis,
|
|
35
|
+
_check_size_consistency,
|
|
35
36
|
)
|
|
36
37
|
from .types import GatherScatterModeStr, IntScalar, PrecisionLike, Scalar
|
|
37
38
|
|
|
@@ -1649,39 +1650,62 @@ def _broadcast_axes(
|
|
|
1649
1650
|
def broadcast_to(
|
|
1650
1651
|
a: NamedOrNumeric, axes: AxisSpec, ensure_order: bool = True, enforce_no_extra_axes: bool = True
|
|
1651
1652
|
) -> NamedArray:
|
|
1652
|
-
"""
|
|
1653
|
-
|
|
1654
|
-
|
|
1655
|
-
|
|
1653
|
+
"""Broadcast ``a`` so that it has the given axes.
|
|
1654
|
+
|
|
1655
|
+
If ``ensure_order`` is ``True`` (default) then the returned array's axes are
|
|
1656
|
+
arranged in the same order as ``axes``. Otherwise existing axes may remain in
|
|
1657
|
+
their current order, though they may still be moved to the front if new axes
|
|
1658
|
+
are added.
|
|
1656
1659
|
|
|
1657
|
-
If enforce_no_extra_axes is True and
|
|
1660
|
+
If ``enforce_no_extra_axes`` is ``True`` and ``a`` has axes that are not in
|
|
1661
|
+
``axes`` then a ``ValueError`` is raised.
|
|
1658
1662
|
"""
|
|
1659
|
-
|
|
1663
|
+
|
|
1664
|
+
axes_dict = axis_spec_to_shape_dict(axes)
|
|
1665
|
+
axes_tuple = axis_spec_to_tuple(axes)
|
|
1660
1666
|
|
|
1661
1667
|
if not isinstance(a, NamedArray):
|
|
1662
1668
|
a = named(jnp.asarray(a), ())
|
|
1663
1669
|
|
|
1664
1670
|
assert isinstance(a, NamedArray) # mypy gets confused
|
|
1665
1671
|
|
|
1666
|
-
|
|
1667
|
-
|
|
1672
|
+
a_axes_dict = axis_spec_to_shape_dict(a.axes)
|
|
1673
|
+
|
|
1674
|
+
# fill in missing sizes and check for mismatches
|
|
1675
|
+
for name, sz in list(axes_dict.items()):
|
|
1676
|
+
if sz is None:
|
|
1677
|
+
if name not in a_axes_dict:
|
|
1678
|
+
raise ValueError(
|
|
1679
|
+
f"Cannot broadcast: size for axis '{name}' is unspecified and it does not exist in array"
|
|
1680
|
+
)
|
|
1681
|
+
axes_dict[name] = a_axes_dict[name]
|
|
1682
|
+
elif name in a_axes_dict:
|
|
1683
|
+
_check_size_consistency(axes, a.axes, name, sz, a_axes_dict[name])
|
|
1668
1684
|
|
|
1669
|
-
|
|
1685
|
+
extra_axis_names = [ax.name for ax in a.axes if ax.name not in axes_dict]
|
|
1686
|
+
if enforce_no_extra_axes and extra_axis_names:
|
|
1687
|
+
raise ValueError(
|
|
1688
|
+
f"Cannot broadcast {a.shape} to {axes_dict}: extra axes present {extra_axis_names}"
|
|
1689
|
+
)
|
|
1670
1690
|
|
|
1671
|
-
|
|
1691
|
+
axes_names_in_a = {ax.name for ax in a.axes}
|
|
1692
|
+
to_add = tuple(
|
|
1693
|
+
Axis(axis_name(ax), axes_dict[axis_name(ax)])
|
|
1694
|
+
for ax in axes_tuple
|
|
1695
|
+
if axis_name(ax) not in axes_names_in_a
|
|
1696
|
+
)
|
|
1672
1697
|
|
|
1673
|
-
|
|
1674
|
-
raise ValueError(f"Cannot broadcast {a.shape} to {axis_spec_to_shape_dict(axes)}: extra axes present")
|
|
1698
|
+
all_axes = to_add + a.axes
|
|
1675
1699
|
|
|
1676
|
-
extra_axes = tuple(ax for ax in a.axes if ax not in
|
|
1700
|
+
extra_axes = tuple(ax for ax in a.axes if ax.name not in axes_dict)
|
|
1677
1701
|
|
|
1678
|
-
# broadcast whatever we need to the front and reorder
|
|
1679
1702
|
a_array = jnp.broadcast_to(a.array, [ax.size for ax in all_axes])
|
|
1680
1703
|
a = NamedArray(a_array, all_axes)
|
|
1681
1704
|
|
|
1682
|
-
|
|
1683
|
-
|
|
1684
|
-
|
|
1705
|
+
axes_tuple_complete = tuple(Axis(axis_name(ax), axes_dict[axis_name(ax)]) for ax in axes_tuple)
|
|
1706
|
+
|
|
1707
|
+
if ensure_order and not _is_subsequence(axes_tuple_complete, all_axes):
|
|
1708
|
+
a = a.rearrange(axes_tuple_complete + extra_axes)
|
|
1685
1709
|
|
|
1686
1710
|
return typing.cast(NamedArray, a)
|
|
1687
1711
|
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev392"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|