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 constructsAxistuples from keyword argumentsconcat_axes– Merges multiple axis specifications while enforcing uniqueness constraintsunion_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/(includingLinear,MLP, and attention mechanisms) acceptAxisobjects as constructor arguments, making models portable across batch sizes and sequence lengths without configuration changes. - Shape Validation: Functions like
unsize_axesandeliminate_axesallow 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
Axisabstraction that transforms JAX arrays into self-documenting tensors with named dimensions, implemented inlib/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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →