haliax 1.4.dev392__tar.gz → 1.4.dev394__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 (117) hide show
  1. {haliax-1.4.dev392 → haliax-1.4.dev394}/AGENTS.md +1 -0
  2. {haliax-1.4.dev392 → haliax-1.4.dev394}/PKG-INFO +1 -1
  3. haliax-1.4.dev394/docs/primer.md +114 -0
  4. {haliax-1.4.dev392 → haliax-1.4.dev394}/mkdocs.yml +2 -0
  5. haliax-1.4.dev394/src/haliax/__about__.py +1 -0
  6. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/core.py +41 -17
  7. haliax-1.4.dev392/src/haliax/__about__.py +0 -1
  8. {haliax-1.4.dev392 → haliax-1.4.dev394}/.coveragerc +0 -0
  9. {haliax-1.4.dev392 → haliax-1.4.dev394}/.flake8 +0 -0
  10. {haliax-1.4.dev392 → haliax-1.4.dev394}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev392 → haliax-1.4.dev394}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev392 → haliax-1.4.dev394}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev392 → haliax-1.4.dev394}/.github/workflows/run_tests.yaml +0 -0
  14. {haliax-1.4.dev392 → haliax-1.4.dev394}/.gitignore +0 -0
  15. {haliax-1.4.dev392 → haliax-1.4.dev394}/.playbooks/add-types.md +0 -0
  16. {haliax-1.4.dev392 → haliax-1.4.dev394}/.playbooks/wrap-non-named.md +0 -0
  17. {haliax-1.4.dev392 → haliax-1.4.dev394}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev392 → haliax-1.4.dev394}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev392 → haliax-1.4.dev394}/CONTRIBUTING.md +0 -0
  20. {haliax-1.4.dev392 → haliax-1.4.dev394}/LICENSE +0 -0
  21. {haliax-1.4.dev392 → haliax-1.4.dev394}/README.md +0 -0
  22. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/api.md +0 -0
  23. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/faq.md +0 -0
  28. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/index.md +0 -0
  40. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/nn.md +0 -0
  43. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/scan.md +0 -0
  48. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/typing.md +0 -0
  51. {haliax-1.4.dev392 → haliax-1.4.dev394}/docs/vmap.md +0 -0
  52. {haliax-1.4.dev392 → haliax-1.4.dev394}/pyproject.toml +0 -0
  53. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/__init__.py +0 -0
  54. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/__init__.py +0 -0
  55. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/compile_utils.py +0 -0
  56. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/dot.py +0 -0
  57. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/einsum.py +0 -0
  58. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/fp8.py +0 -0
  59. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/parsing.py +0 -0
  60. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/rearrange.py +0 -0
  61. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/scan.py +0 -0
  62. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/debug.py +0 -0
  66. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/haxtyping.py +0 -0
  67. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/hof.py +0 -0
  68. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/jax_utils.py +0 -0
  69. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/__init__.py +0 -0
  70. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/activations.py +0 -0
  71. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/attention.py +0 -0
  72. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/conv.py +0 -0
  73. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/dropout.py +0 -0
  74. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/embedding.py +0 -0
  75. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/linear.py +0 -0
  76. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/loss.py +0 -0
  77. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/mlp.py +0 -0
  78. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/normalization.py +0 -0
  79. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/pool.py +0 -0
  80. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/nn/scan.py +0 -0
  81. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/ops.py +0 -0
  82. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev392 → haliax-1.4.dev394}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/core_test.py +0 -0
  92. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_attention.py +0 -0
  93. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_axis.py +0 -0
  94. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_conv.py +0 -0
  95. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_debug.py +0 -0
  96. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_dot.py +0 -0
  97. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_dtype_typing.py +0 -0
  98. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_einsum.py +0 -0
  99. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_fp8.py +0 -0
  100. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_hof.py +0 -0
  101. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_int8.py +0 -0
  102. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_namedarray_typing.py +0 -0
  103. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_nn.py +0 -0
  104. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_ops.py +0 -0
  105. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_parsing.py +0 -0
  106. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_partitioning.py +0 -0
  107. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_pool.py +0 -0
  108. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_random.py +0 -0
  109. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_rearrange.py +0 -0
  110. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_scan.py +0 -0
  111. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_scatter_gather.py +0 -0
  112. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_specialized_fns.py +0 -0
  113. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_state_dict.py +0 -0
  114. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_tree_util.py +0 -0
  115. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_utils.py +0 -0
  116. {haliax-1.4.dev392 → haliax-1.4.dev394}/tests/test_visualize_sharding.py +0 -0
  117. {haliax-1.4.dev392 → haliax-1.4.dev394}/uv.lock +0 -0
@@ -74,3 +74,4 @@ repository. Follow these notes when implementing new features or fixing bugs.
74
74
 
75
75
  ## Documentation
76
76
  - Public functions and modules require docstrings. If behavior is non‑obvious, add examples in `docs/`.
77
+ - For a concise overview of Haliax aimed at LLM agents, see [docs/primer.md](docs/primer.md).
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev392
3
+ Version: 1.4.dev394
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,114 @@
1
+ # Haliax Primer
2
+
3
+ Haliax provides named tensors built on top of JAX. This primer is written for LLM agents and other downstream libraries and collects the core ideas for quick reference.
4
+
5
+ ## Axes and Named Arrays
6
+
7
+ Arrays are indexed by `Axis` objects. You can define them explicitly or generate several with `make_axes`.
8
+ You may also specify shapes with a **shape dict**, mapping axis names to sizes.
9
+
10
+ ```python
11
+ import haliax as hax
12
+ from haliax import Axis
13
+
14
+ Batch = Axis("batch", 4)
15
+ Feature = Axis("feature", 8)
16
+ # or: Batch, Feature = hax.make_axes(batch=4, feature=8)
17
+ # using Axis objects
18
+ x = hax.zeros((Batch, Feature))
19
+ # or using a shape dict
20
+ shape = {"batch": 4, "feature": 8}
21
+ x = hax.zeros(shape)
22
+ ```
23
+
24
+ Most functions accept either axes or shape dicts interchangeably.
25
+
26
+ A tensor with named axes is a [`NamedArray`][haliax.NamedArray]. Elementwise operations mirror `jax.numpy` but accept named axes.
27
+
28
+ ## Indexing and Broadcasting
29
+
30
+ Use axis names when slicing. Dictionaries are convenient for several axes:
31
+
32
+ ```python
33
+ first = x["batch", 0]
34
+ sub = x["batch", 1:3]
35
+ # or with a dict
36
+ first = x[{"batch": 0}]
37
+ sub = x[{"batch": slice(1, 3)}]
38
+ ```
39
+
40
+ Axes broadcast by matching names. `broadcast_axis` adds a new axis to an array:
41
+
42
+ ```python
43
+ row = hax.arange(Feature)
44
+ outer = row.broadcast_axis(Batch) * hax.arange(Batch)
45
+ ```
46
+
47
+ See [Indexing and Slicing](indexing.md) and [Broadcasting](broadcasting.md) for details.
48
+
49
+ ## Rearranging Axes
50
+
51
+ `rearrange` changes axis order and can merge or split axes using einops‑style syntax. It is useful when interfacing with positional APIs.
52
+
53
+ ```python
54
+ # transpose features and batch
55
+ x_t = hax.rearrange(x, "batch feature -> feature batch")
56
+ ```
57
+
58
+ More examples appear in [Rearrange](rearrange.md).
59
+
60
+ ## Matrix Multiplication
61
+
62
+ `dot` contracts over named axes while preserving order independence.
63
+
64
+ ```python
65
+ Weight = Axis("weight", 8)
66
+ w = hax.ones((Feature, Weight))
67
+ prod = hax.dot(x, w, axis=Feature)
68
+ ```
69
+
70
+ For more complex contractions use [`einsum`][haliax.einsum]. See [Matrix Multiplication](matmul.md).
71
+
72
+ ## Scans and Folds
73
+
74
+ Use [`scan`][haliax.scan] or [`fold`][haliax.fold] to apply a function along an axis with optional gradient checkpointing.
75
+
76
+ ```python
77
+ Time = Axis("time", 10)
78
+ sequence = hax.ones((Time, Feature))
79
+
80
+ def add(prev, cur):
81
+ return prev + cur
82
+
83
+ result = hax.fold(add, Time)(hax.zeros((Feature,)), sequence)
84
+ ```
85
+
86
+ See [Scan and Fold](scan.md) for checkpointing policies and stacked modules.
87
+
88
+ ## Partitioning
89
+
90
+ Arrays and modules can be distributed across devices by mapping named axes to mesh axes:
91
+
92
+ ```python
93
+ with hax.axis_mapping({"batch": "data"}):
94
+ sharded = hax.shard(x)
95
+ ```
96
+
97
+ The [Partitioning](partitioning.md) guide explains how to set up device meshes and shard arrays.
98
+
99
+ ## Typing Support
100
+
101
+ Type annotations use `haliax.haxtyping` which extends `jaxtyping`:
102
+
103
+ ```python
104
+ import haliax.haxtyping as ht
105
+
106
+ def f(t: ht.Float[hax.NamedArray, "batch feature"]):
107
+ ...
108
+ ```
109
+
110
+ See [Typing](typing.md) for matching runtime checks and dtype-aware annotations.
111
+
112
+ ---
113
+
114
+ This primer highlights common patterns. The [cheatsheet](cheatsheet.md) lists many additional conversions from JAX to Haliax.
@@ -100,3 +100,5 @@ nav:
100
100
  - Serialization: 'state-dict.md'
101
101
  - API Reference: 'api.md'
102
102
  - FAQ: 'faq.md'
103
+ - LLMs:
104
+ - "LLM Primer": 'primer.md'
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev394"
@@ -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