# Haliax for Named Array Programming in JAX: A Complete Guide

> Discover Haliax for named array programming in JAX. Replace positional dimensions with named axes to eliminate shape mismatch bugs and write self-documenting code.

- Repository: [The Marin Project/marin](https://github.com/marin-community/marin)
- Tags: deep-dive
- Published: 2026-09-10

---

**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**:

```python

# 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**:

```python
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)](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/axis.py):

```python
@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`](https://github.com/marin-community/marin/blob/main/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)](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/core.py)) wraps a standard JAX array together with its axis specification:

```python
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`:

```python
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:

```python
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)](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 axes
- **`einsum`** – Einstein summation using axis names instead of subscripts
- **`rearrange`** – 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`](https://github.com/marin-community/marin/blob/main/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)](https://github.com/marin-community/marin/blob/main/lib/haliax/src/haliax/partitioning.py) maps named axes to **SPMD mesh dimensions**:

```python
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

```python
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 `Axis` dataclass and `NamedArray` wrapper, located in `lib/haliax/src/haliax/`.
- **Axis objects carry `name` and `size`**, 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.
- **`dslice` enables JIT-compatible dynamic slicing** by separating static lengths from dynamic indices.
- **Higher-level operations** in [`ops.py`](https://github.com/marin-community/marin/blob/main/ops.py) and distributed support in [`partitioning.py`](https://github.com/marin-community/marin/blob/main/partitioning.py) extend 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)](https://github.com/marin-community/marin/blob/main/lib/haliax/docs/tutorial.md) with executable Colab notebooks. The test suite in [`lib/haliax/tests/`](https://github.com/marin-community/marin/tree/main/lib/haliax/tests) provides additional usage patterns for axis manipulation and distributed partitioning.