Haliax for Named Array Programming in JAX: A Complete Guide
Haliax is a JAX library that replaces positional array dimensions with semantic Axis names, eliminating shape-mismatch bugs and making code self-documenting.
Haliax ships as part of the Marin ecosystem and provides a named-array abstraction layer over standard JAX arrays. Rather than tracking dimensions by position, you attach explicit Axis objects to arrays—enabling safer indexing, clearer function signatures, and automatic shape validation. This article examines how Haliax for named array programming in JAX works, based on the implementation in marin-community/marin.
What Is Named Array Programming?
Traditional JAX code relies on positional dimensions:
# Which dimension is which? Easy to get wrong.
x: jax.Array # shape (8, 10, 64)
y = jnp.mean(x, axis=1) # averages over... time? features?
Named array programming replaces anonymous positions with explicit axis labels:
import haliax as hax
# Axes carry both name and size
Batch, Time, Features = hax.make_axes(batch=8, time=10, features=64)
data = hax.named(x, (Batch, Time, Features))
# Operations reference axes by name, not position
y = hax.mean(data, axis=Time) # unambiguous: reduces over temporal axis
The core insight, as implemented in lib/haliax/src/haliax/, is that names prevent an entire class of shape bugs while making tensor operations self-describing.
The Axis Dataclass: Foundation of Haliax
All naming in Haliax flows through the Axis dataclass, defined in [lib/haliax/src/haliax/axis.py](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/axis.py):
@dataclass(frozen=True)
class Axis:
name: str
size: int
Axis objects are immutable and hashable, so they can be used as dictionary keys and function defaults. The size attribute holds the static dimension length, while name provides the semantic identifier.
Creating and Managing Axes
Haliax provides several utilities for axis construction and manipulation, all residing in axis.py:
| Function | Purpose |
|---|---|
make_axes(**kwargs) |
Create multiple Axis objects from keyword arguments |
axis_spec_to_shape_dict(spec) |
Convert any axis specification to {name: size} mapping |
concat_axes(*specs) |
Combine axis specifications sequentially |
union_axes(*specs) |
Merge specifications, checking for consistent sizes |
eliminate_axes(spec, *to_remove) |
Remove specified axes from a specification |
replace_axis(spec, old, new) |
Substitute one axis for another |
These utilities ensure that axis manipulation remains type-safe—attempting to concatenate mismatched sizes or reference non-existent axes raises clear errors at trace time.
NamedArray: JAX Arrays with Attached Semantics
The NamedArray class (in [lib/haliax/src/haliax/core.py](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/core.py)) wraps a standard JAX array together with its axis specification:
import jax.numpy as jnp
# Create a NamedArray from raw JAX array and axes
raw = jnp.ones((8, 10, 64))
named = hax.named(raw, (Batch, Time, Features))
# Underlying data is still a JAX array
assert isinstance(named.array, jnp.ndarray)
# But axes are tracked explicitly
assert named.axes == (Batch, Time, Features)
NamedArray supports all standard JAX operations while preserving axis information. The to_jax_shape() method (and related utilities) translate named shapes back to positional tuples when interfacing with raw JAX code.
Dynamic Slicing with dslice
A notable Haliax feature is JIT-compatible dynamic slicing, implemented via the dslice class (lines 89–124 in the source). Unlike standard Python slices, dslice separates static length from dynamic start index, enabling flexible indexing inside jax.jit:
from haliax import dslice
def extract_window(x: hax.NamedArray, start: int) -> hax.NamedArray:
# dslice(start, length) — length is static, start is dynamic
return x[:, dslice(start, 5), :] # 5-step window at dynamic position
# JIT compilation works despite dynamic start
windowed = jax.jit(extract_window)(data, start=3)
This pattern is essential for variable-length sequences and sliding window operations where the slice position depends on runtime values.
Axis Reordering and Partial Orders
Haliax handles dimension permutation through rearrange_for_partial_order, which supports ellipsis notation for flexible reordering:
from haliax import Ellipsis
# Move Batch axis to the end, keep others in relative order
new_order = hax.rearrange_for_partial_order(
(Ellipsis, Batch), # target pattern: [*other_axes, Batch]
data.axes # original axis specification
)
# Result: (Time, Features, Batch)
The ellipsis (... or hax.Ellipsis) acts as a catch-all placeholder, making patterns robust to changes in the full axis set. This utility appears in distributed computing contexts where specific axes must align with hardware mesh dimensions.
Higher-Level Operations in Haliax
The [lib/haliax/src/haliax/ops.py](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/ops.py) module extends Haliax with named-aware linear algebra:
dot– matrix multiplication with explicit contraction axeseinsum– Einstein summation using axis names instead of subscriptsrearrange– generalized tensor transposition with name-based patterns
These operations automatically validate axis compatibility and generate optimal JAX primitives under the hood.
Distributed Execution with partitioning.py
For multi-device JAX programs, [lib/haliax/src/haliax/partitioning.py](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/partitioning.py) maps named axes to SPMD mesh dimensions:
from haliax.partitioning import ResourceAxis, with_sharding_constraint
# Declare which logical axes map to which hardware resources
device_mesh = jax.make_mesh((4, 2), (ResourceAxis("data"), ResourceAxis("model")))
# Annotate arrays with sharding constraints using familiar names
sharded = with_sharding_constraint(data, (ResourceAxis("data"), None, None))
Named axes make parallelism specifications readable and less error-prone than positional sharding strings.
Practical Example: End-to-End Workflow
import jax
import jax.numpy as jnp
import haliax as hax
# 1. Define semantic axes
Batch, SeqLen, Embed, Heads, HeadDim = hax.make_axes(
batch=32, seq=512, embed=768, heads=12, head_dim=64
)
# 2. Create named parameters
def init_attention():
w_q = hax.named(jax.random.normal(k1, (Embed.size, Heads.size, HeadDim.size)),
(Embed, Heads, HeadDim))
w_k = hax.named(jax.random.normal(k2, (Embed.size, Heads.size, HeadDim.size)),
(Embed, Heads, HeadDim))
return w_q, w_k
# 3. Write named computation (self-attention snippet)
def attention_step(q: hax.NamedArray, k: hax.NamedArray, v: hax.NamedArray):
# q: (Batch, SeqLen, Heads, HeadDim)
scores = hax.dot("HeadDim", q, k) # contracts over HeadDim
scores = scores / jnp.sqrt(HeadDim.size)
attn_weights = hax.softmax(scores, axis=SeqLen)
return hax.dot("SeqLen", attn_weights, v) # (Batch, Heads, HeadDim)
# 4. JIT compile — names are resolved to shapes at trace time
batched_attention = jax.vmap(attention_step, in_axes=(Batch, Batch, Batch))
compiled = jax.jit(batched_attention)
Summary
- Haliax provides named-axis programming for JAX through the
Axisdataclass andNamedArraywrapper, located inlib/haliax/src/haliax/. - Axis objects carry
nameandsize, enabling self-documenting array code that catches shape mismatches at compile time. - Core utilities (
make_axes,concat_axes,eliminate_axes,rearrange_for_partial_order) manipulate axis specifications safely. dsliceenables JIT-compatible dynamic slicing by separating static lengths from dynamic indices.- Higher-level operations in
ops.pyand distributed support inpartitioning.pyextend Haliax to full ML workloads.
Frequently Asked Questions
How does Haliax differ from standard JAX?
Standard JAX uses positional dimensions where axis=0 and axis=1 are indistinguishable without documentation. Haliax attaches semantic names to every dimension, making axis=Batch and axis=Time self-describing. This prevents bugs from transposed tensors and makes function contracts explicit.
Can I use Haliax with existing JAX code?
Yes. NamedArray wraps ordinary JAX arrays, and to_jax_shape() converts named specifications back to positional tuples. You can extract the underlying array with .array when interfacing with libraries that expect raw JAX arrays, then re-wrap results with hax.named().
What is dslice and when do I need it?
dslice is a JIT-compatible dynamic slice that separates static length from dynamic start position. Use it when slice bounds depend on runtime values—for example, extracting variable-length subsequences in language models or sliding windows in time-series analysis. Standard Python slices fail inside jax.jit when bounds are not static.
Where can I learn more about Haliax usage?
The Marin repository includes interactive tutorials at [lib/haliax/docs/tutorial.md](https://github.com/marin-community/marin/blob/main/lib/haliax/docs/tutorial.md) with executable Colab notebooks. The test suite in lib/haliax/tests/ provides additional usage patterns for axis manipulation and distributed partitioning.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →