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

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 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 and distributed support in 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) 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →