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.
- {haliax-1.4.dev393 → haliax-1.4.dev394}/AGENTS.md +1 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/PKG-INFO +1 -1
- haliax-1.4.dev394/docs/primer.md +114 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/mkdocs.yml +2 -0
- haliax-1.4.dev394/src/haliax/__about__.py +1 -0
- haliax-1.4.dev393/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.coveragerc +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.flake8 +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.gitignore +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/LICENSE +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/README.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/api.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/css/material.css +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/faq.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/fp8.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/index.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/indexing.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/matmul.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/nn.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/partitioning.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/rearrange.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/requirements.txt +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/scan.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/state-dict.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/tutorial.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/typing.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/docs/vmap.md +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/pyproject.toml +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/core.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/random.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/types.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/util.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/core_test.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_attention.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_axis.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_conv.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_debug.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_dot.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_hof.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_int8.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_nn.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_ops.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_pool.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_random.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_scan.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_utils.py +0 -0
- {haliax-1.4.dev393 → haliax-1.4.dev394}/tests/test_visualize_sharding.py +0 -0
- {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.
|
|
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.
|
|
@@ -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
|
|
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
|