haliax 1.4.dev391__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.
Files changed (116) hide show
  1. {haliax-1.4.dev391 → haliax-1.4.dev393}/AGENTS.md +6 -8
  2. {haliax-1.4.dev391 → haliax-1.4.dev393}/PKG-INFO +1 -1
  3. haliax-1.4.dev393/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/core.py +41 -17
  5. haliax-1.4.dev391/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev391 → haliax-1.4.dev393}/.coveragerc +0 -0
  7. {haliax-1.4.dev391 → haliax-1.4.dev393}/.flake8 +0 -0
  8. {haliax-1.4.dev391 → haliax-1.4.dev393}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev391 → haliax-1.4.dev393}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev391 → haliax-1.4.dev393}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev391 → haliax-1.4.dev393}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev391 → haliax-1.4.dev393}/.gitignore +0 -0
  13. {haliax-1.4.dev391 → haliax-1.4.dev393}/.playbooks/add-types.md +0 -0
  14. {haliax-1.4.dev391 → haliax-1.4.dev393}/.playbooks/wrap-non-named.md +0 -0
  15. {haliax-1.4.dev391 → haliax-1.4.dev393}/.pre-commit-config.yaml +0 -0
  16. {haliax-1.4.dev391 → haliax-1.4.dev393}/.readthedocs.yaml +0 -0
  17. {haliax-1.4.dev391 → haliax-1.4.dev393}/CONTRIBUTING.md +0 -0
  18. {haliax-1.4.dev391 → haliax-1.4.dev393}/LICENSE +0 -0
  19. {haliax-1.4.dev391 → haliax-1.4.dev393}/README.md +0 -0
  20. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/api.md +0 -0
  21. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/broadcasting.md +0 -0
  22. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/cheatsheet.md +0 -0
  23. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/css/material.css +0 -0
  24. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/css/mkdocstrings.css +0 -0
  25. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/faq.md +0 -0
  26. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/data_parallel_mesh.png +0 -0
  27. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  28. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_1d.png +0 -0
  29. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_1d_zero.png +0 -0
  30. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d.png +0 -0
  31. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  32. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  33. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  34. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  35. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/figures/device_mesh_2d_zero.png +0 -0
  36. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/fp8.md +0 -0
  37. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/index.md +0 -0
  38. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/indexing.md +0 -0
  39. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/matmul.md +0 -0
  40. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/nn.md +0 -0
  41. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/partitioning.md +0 -0
  42. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/rearrange.ipynb +0 -0
  43. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/rearrange.md +0 -0
  44. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/requirements.txt +0 -0
  45. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/scan.md +0 -0
  46. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/state-dict.md +0 -0
  47. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/tutorial.md +0 -0
  48. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/typing.md +0 -0
  49. {haliax-1.4.dev391 → haliax-1.4.dev393}/docs/vmap.md +0 -0
  50. {haliax-1.4.dev391 → haliax-1.4.dev393}/mkdocs.yml +0 -0
  51. {haliax-1.4.dev391 → haliax-1.4.dev393}/pyproject.toml +0 -0
  52. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/__init__.py +0 -0
  53. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/__init__.py +0 -0
  54. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/compile_utils.py +0 -0
  55. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/dot.py +0 -0
  56. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/einsum.py +0 -0
  57. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/fp8.py +0 -0
  58. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/parsing.py +0 -0
  59. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/rearrange.py +0 -0
  60. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/scan.py +0 -0
  61. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/state_dict.py +0 -0
  62. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/_src/util.py +0 -0
  63. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/axis.py +0 -0
  64. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/debug.py +0 -0
  65. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/haxtyping.py +0 -0
  66. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/hof.py +0 -0
  67. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/jax_utils.py +0 -0
  68. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/__init__.py +0 -0
  69. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/activations.py +0 -0
  70. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/attention.py +0 -0
  71. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/conv.py +0 -0
  72. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/dropout.py +0 -0
  73. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/embedding.py +0 -0
  74. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/linear.py +0 -0
  75. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/loss.py +0 -0
  76. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/mlp.py +0 -0
  77. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/normalization.py +0 -0
  78. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/pool.py +0 -0
  79. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/nn/scan.py +0 -0
  80. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/ops.py +0 -0
  81. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/partitioning.py +0 -0
  82. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/quantization.py +0 -0
  83. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/random.py +0 -0
  84. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/specialized_fns.py +0 -0
  85. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/state_dict.py +0 -0
  86. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/tree_util.py +0 -0
  87. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/types.py +0 -0
  88. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/util.py +0 -0
  89. {haliax-1.4.dev391 → haliax-1.4.dev393}/src/haliax/wrap.py +0 -0
  90. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/core_test.py +0 -0
  91. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_attention.py +0 -0
  92. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_axis.py +0 -0
  93. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_dtype_typing.py +0 -0
  97. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_einsum.py +0 -0
  98. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_fp8.py +0 -0
  99. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_hof.py +0 -0
  100. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_int8.py +0 -0
  101. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_namedarray_typing.py +0 -0
  102. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_nn.py +0 -0
  103. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_ops.py +0 -0
  104. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_parsing.py +0 -0
  105. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_partitioning.py +0 -0
  106. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_pool.py +0 -0
  107. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_random.py +0 -0
  108. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_rearrange.py +0 -0
  109. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_scan.py +0 -0
  110. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_scatter_gather.py +0 -0
  111. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_specialized_fns.py +0 -0
  112. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_state_dict.py +0 -0
  113. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_tree_util.py +0 -0
  114. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_utils.py +0 -0
  115. {haliax-1.4.dev391 → haliax-1.4.dev393}/tests/test_visualize_sharding.py +0 -0
  116. {haliax-1.4.dev391 → haliax-1.4.dev393}/uv.lock +0 -0
@@ -17,7 +17,7 @@ repository. Follow these notes when implementing new features or fixing bugs.
17
17
  ## Playbook
18
18
 
19
19
  - Adding Haliax-style tensor typing annotations are described in @.playbooks/add-types.md
20
- - Wrapping standard JAX functions so they operate on `NamedArray` is explained in @.playbooks/wrap-non-named.md
20
+ - [Wrapping standard JAX functions](.playbooks/wrap-non-named.md) so they operate on `NamedArray`
21
21
 
22
22
  ## Code Style
23
23
 
@@ -58,15 +58,14 @@ repository. Follow these notes when implementing new features or fixing bugs.
58
58
 
59
59
  * **Generic code**: many utilities are written with Python generics and dataclasses. Where possible,
60
60
  write reusable functions or classes that operate over TypeVars instead of hard coding concrete types.
61
- * **Configurations**: configuration files are dataclasses loaded via `draccus`. Keep configs
62
- declarative and typed.
63
- * **Reproducibility**: Levanter aims for deterministic training where possible. Avoid sources of
61
+ * **Reproducibility**: Haliax aims for determinism where possible. Avoid sources of
64
62
  nondeterminism unless explicitly required.
65
63
  * Prefer Stacked with fold or scan over writing custom loops, for better compile times and gradient checkpointing support
64
+ * For configuration, we prefer frozen dataclasses over dictionaries.
66
65
 
67
66
  ## Library conventions
68
- - Haliax revolves around `NamedArray` and explicit `Axis` objects. Prefer APIs that accept
69
- axes or axis names rather than hard‑coding positional dimensions.
67
+ - Haliax revolves around `NamedArray` and named shapes, either via Axis objects or "shape dicts" (e.g. `{"batch": 42, "embed": 16}).
68
+ Prefer APIs that accept axes or axis names rather than hard‑coding positional dimensions. In particular, use AxisSpec and AxisSelection where possible.
70
69
  - Utilities should be written so they work with arbitrary axis names. Avoid relying on
71
70
  fixed axis orders when possible.
72
71
  - Use the provided modules in `haliax.nn` or Equinox when building neural network layers.
@@ -74,5 +73,4 @@ repository. Follow these notes when implementing new features or fixing bugs.
74
73
  for a float32 array with a "batch" axis, or `ht.Float[NamedArray, "batch"]` for any floating point dtype.
75
74
 
76
75
  ## Documentation
77
- - Public functions and modules require docstrings. If behavior is non‑obvious,
78
- add examples in `docs/`.
76
+ - Public functions and modules require docstrings. If behavior is non‑obvious, add examples in `docs/`.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev391
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
- Broadcasts a so that it has the given axes.
1654
- If ensure_order is True (default), then the returned array will have the same axes in the same order as the given
1655
- axes. Otherwise, the axes may not be moved if they are already in the array. The axes may not be contiguous however
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 the array has axes that are not in axes, then a ValueError is raised.
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
- axes = axis_spec_to_tuple(axes)
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
- if a.axes == axes:
1667
- return a
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
- to_add = tuple(ax for ax in axes if ax not in a.axes)
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
- all_axes = to_add + a.axes
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
- if enforce_no_extra_axes and len(all_axes) != len(axes):
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 axes)
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
- # if the new axes are already in the right order, then we're done
1683
- if ensure_order and not _is_subsequence(axes, all_axes):
1684
- a = a.rearrange(axes + extra_axes)
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.dev391"
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes