Understanding Haliax's Role as the Core Tensor/Named Axis Library in Marin

Haliax serves as the foundational tensor abstraction layer that enables Marin to perform all tensor operations using explicit named axes rather than anonymous integer indices.

The marin-community/marin repository relies on Haliax to provide the semantic backbone for its machine learning pipelines, transforming raw JAX arrays into self-documenting, composable data structures. By treating axes as first-class objects with human-readable names and explicit sizes, Haliax eliminates shape-mismatch bugs and streamlines distributed training across complex hardware meshes.

Core Responsibilities of Haliax

Haliax is not merely a utility library—it is the conceptual foundation upon which Marin's training, evaluation, and inference systems are built. Its responsibilities span four critical domains.

Named Axis Semantics

At the heart of Haliax lies the Axis dataclass defined in lib/haliax/src/haliax/axis.py (lines 18–34). Unlike conventional tensor libraries that track shapes via tuples of integers, Haliax requires every tensor dimension to be associated with an Axis object containing a name and size. This design choice makes shape contracts explicit and machine-verifiable throughout Marin's codebase.

Axis Specification Utilities

Haliax provides a comprehensive toolbox for manipulating axis collections without resorting to error-prone integer indexing. Key functions include:

  • make_axes – Programmatically constructs Axis tuples from keyword arguments
  • concat_axes – Merges multiple axis specifications while enforcing uniqueness constraints
  • union_axes – Combines axis sets, dropping duplicates (with strict size-matching validation)
  • replace_axis – Substitutes specific axes within a specification while preserving order

These utilities appear throughout Marin's training scripts in lib/marin/src/marin/experiment/train.py and model loading logic in lib/marin/src/marin/evaluation/model_loading.py.

Distributed Partitioning and Mesh Management

The haliax.partitioning module (located in lib/haliax/src/haliax/partitioning.py) bridges the gap between logical axis names and physical hardware layouts. Through functions like set_mesh and axis_mapping, Haliax translates abstract Axis descriptions into concrete JAX device meshes.

Marin's distributed training engine (Levanter) consumes these mappings to implement data-parallel and tensor-parallel strategies. The ResourceAxis hierarchy—distinguishing between DATA and MODEL parallel resources—enables sophisticated sharding without hard-coding device counts into model definitions.

High-Level Tensor Constructors

Haliax wraps JAX's array constructors to automatically attach axis metadata. The haliax.named function (implemented in lib/haliax/src/haliax/core.py) creates jax.Array instances bound to specific Axis tuples, while convenience functions like haliax.zeros, haliax.ones, and haliax.arange provide drop-in replacements for jax.numpy operations that preserve semantic axis information.

Implementation Examples

The following patterns demonstrate how Marin leverages Haliax for safe, readable tensor manipulation.

Creating Named Tensors

import haliax as hax
import jax.numpy as jnp

# Define semantic axes

batch, seq = hax.make_axes(batch=8, seq=128)

# Create tensor with attached metadata

x = hax.named(jnp.zeros((8, 128)), (batch, seq))

print(x.shape)  # (8, 128)

print(x.axes)   # (batch, seq)

This code instantiates an array where the first dimension is explicitly labeled "batch" (size 8) and the second "seq" (size 128), preventing accidental transposition errors during model forward passes.

Manipulating Axis Specifications

import haliax as hax

# Combine specifications safely

spec1 = {"batch": 8, "seq": 128}
spec2 = hax.Axis("head", 12)

# concat_axes raises ValueError on name collisions

combined = hax.concat_axes(spec1, spec2)

# Result: {"batch": 8, "seq": 128, "head": 12}

# union_axes deduplicates with size validation

unified = hax.union_axes(combined, {"batch": 8})

# Result: {"batch": 8, "seq": 128, "head": 12}

These functions are implemented in lib/haliax/src/haliax/axis.py (lines 98–156) and provide runtime guarantees that axis names remain unique and sizes remain consistent across tensor operations.

Configuring Distributed Meshes

import haliax as hax
from haliax.partitioning import set_mesh, ResourceAxis

# Map logical axes to parallelization strategies

mesh = {
    "batch": ResourceAxis.DATA,
    "head": ResourceAxis.MODEL
}
set_mesh(mesh)  # Registers globally for JAX compilation

This configuration, used by Marin's Levanter integration, assigns the "batch" axis to data parallelism and "head" (attention heads) to model parallelism, allowing the same model code to run efficiently on different hardware topologies.

Safe Reshaping with Partial Orders

import haliax as hax

axes = (
    hax.Axis("batch", 8),
    hax.Axis("seq", 128),
    hax.Axis("head", 12)
)

# Enforce "head" comes first, others follow arbitrarily

partial = (hax.Axis("head", ...), ...)
reordered = hax.rearrange_for_partial_order(partial, axes)

# Returns: (head, batch, seq)

The rearrange_for_partial_order function (lines 445–514 in axis.py) implements a greedy constraint-satisfaction algorithm that respects programmer-specified axis ordering requirements while automatically determining positions for unspecified dimensions.

Integration with Marin's Architecture

Haliax's named-axis paradigm permeates Marin's architecture from data loading through model definition:

  • Model Definitions: Neural network layers in lib/haliax/src/haliax/nn/ (including Linear, MLP, and attention mechanisms) accept Axis objects as constructor arguments, making models portable across batch sizes and sequence lengths without configuration changes.
  • Shape Validation: Functions like unsize_axes and eliminate_axes allow Marin to safely broadcast or reduce dimensions by name rather than position, catching shape mismatches at graph-compile time rather than runtime.
  • Checkpointing: Axis metadata travels with serialized tensors, ensuring that models trained on one hardware topology can be resumed on another without manual shape reconciliation.

Summary

  • Haliax provides the foundational Axis abstraction that transforms JAX arrays into self-documenting tensors with named dimensions, implemented in lib/haliax/src/haliax/axis.py.
  • Axis manipulation utilities (concat_axes, union_axes, make_axes) enable declarative shape management throughout Marin's training and evaluation pipelines.
  • Partitioning infrastructure (set_mesh, ResourceAxis) maps logical axis names to physical device meshes, powering Marin's distributed training capabilities via Levanter.
  • High-level constructors (haliax.named, haliax.zeros) ensure that all tensors in Marin carry semantic metadata from creation, preventing silent shape errors.

Frequently Asked Questions

What makes Haliax different from standard JAX array handling?

Standard JAX represents tensor shapes as anonymous tuples of integers (e.g., (8, 128, 64)), requiring developers to track which dimension corresponds to batch, sequence, or features manually. Haliax replaces these integers with Axis objects containing explicit name and size fields, making tensor shapes self-describing and enabling compile-time verification of shape compatibility across operations.

How does Haliax support distributed training in Marin?

Haliax's partitioning module allows developers to map logical axis names (like "batch" or "head") to physical parallelism strategies (ResourceAxis.DATA or ResourceAxis.MODEL) using set_mesh. This separation of concerns enables Marin's training engine to automatically shard tensors across available hardware without modifying model code, as the same Axis objects guide both computation and distribution.

Where are the core Haliax components located in the Marin repository?

The fundamental components reside in lib/haliax/src/haliax/: axis.py contains the Axis dataclass and specification utilities; core.py implements high-level tensor constructors like haliax.named; partitioning.py manages distributed mesh configuration; and the nn/ subdirectory provides named-axis-aware neural network primitives.

Can Haliax axes be used with existing JAX libraries?

Yes. Haliax tensors are thin wrappers around jax.Array objects that preserve the underlying data buffer while attaching axis metadata. You can extract the raw array via .array or similar accessors when interfacing with libraries that require standard JAX inputs, then re-wrap the result using haliax.named to restore semantic information before returning to Marin code.

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 →