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