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