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

> Explore Haliax, Marin's core tensor and named axis library. Learn how it simplifies tensor operations with explicit named axes for clearer, more robust code.

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

---

**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`](https://github.com/marin-community/marin/blob/main/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`](https://github.com/marin-community/marin/blob/main/lib/marin/src/marin/experiment/train.py) and model loading logic in [`lib/marin/src/marin/evaluation/model_loading.py`](https://github.com/marin-community/marin/blob/main/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`](https://github.com/marin-community/marin/blob/main/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`](https://github.com/marin-community/marin/blob/main/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

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

```python
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`](https://github.com/marin-community/marin/blob/main/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

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

```python
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`](https://github.com/marin-community/marin/blob/main/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`](https://github.com/marin-community/marin/blob/main/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`](https://github.com/marin-community/marin/blob/main/axis.py) contains the `Axis` dataclass and specification utilities; [`core.py`](https://github.com/marin-community/marin/blob/main/core.py) implements high-level tensor constructors like `haliax.named`; [`partitioning.py`](https://github.com/marin-community/marin/blob/main/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.