How Levanter Compiles and Scales JAX Training Across TPU Pods in Marin

Levanter automatically detects TPU topology, constructs a global device mesh, and compiles single-program multi-data (SPMD) XLA executables that distribute training across thousands of TPU cores with automatic synchronization and optional kernel optimizations.

Levanter is the high-performance JAX training library at the core of the Marin repository. It transforms single-device JAX programs into distributed, multi-pod training jobs that scale efficiently across large TPU slices. By automating mesh construction, XLA compilation, and cross-host synchronization, Levanter eliminates the boilerplate typically required for massive-scale training.

Detecting TPU Topology and Hardware Layout

Before compilation begins, Levanter must understand the physical layout of the available compute. The library queries the Cloud TPU metadata service to determine how devices are arranged within the pod.

Reading TPU Metadata in hardware_topology.py

The levanter.utils.hardware_topology module reads the TPU slice layout—represented as AxBxC or SxAxBxC dimensions—from environment metadata. It translates this physical arrangement into a logical Mesh object that describes the grid of available devices.

According to the source code in [levanter/utils/hardware_topology.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/utils/hardware_topology.py), this detection happens automatically when the trainer initializes, requiring no manual configuration for standard TPU topologies.

Constructing the Global Device Mesh

Once the topology is known, Levanter builds a global view of all participating devices. This mesh becomes the foundation for all parallel computation strategies.

Building jax.sharding.Mesh in distributed.py

The levanter.distributed module constructs a jax.sharding.Mesh from the topology information detected in the previous step. This mesh is advertised to the rest of the library via a thread-local context, allowing every JAX primitive to query the current device layout.

As implemented in [levanter/distributed.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/distributed.py), the get_mesh() function retrieves this global mesh and supplies it to the trainer:


# inside levanter/trainer.py – how the mesh is created

from levanter.distributed import get_mesh
from levanter.trainer_state import TrainerState

class Trainer:
    def __init__(self, cfg):
        self.mesh = get_mesh(cfg.mesh)   # reads hardware_topology → Mesh

        self.state = TrainerState(mesh=self.mesh, cfg=cfg)
        # all jax.pjit calls in the trainer automatically use self.mesh

Configuration and Compilation Strategy

Levanter centralizes parallelism strategy and XLA compilation settings in dedicated configuration objects. These determine how the single-device code transforms into a distributed program.

MeshConfig and CompilationConfig

The levanter.config module defines MeshConfig and CompilationConfig objects that store:

  • The device mesh from hardware detection
  • Parallelism strategy (data-parallel, model-parallel, or hybrid)
  • Compilation flags such as jax_disable_most_optimizations

These configurations reside in [levanter/config.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/config.py) and are passed to the Trainer during instantiation.

pjit and XLA Compilation

Core training loops use jax.pjit (or jax.experimental.xmap for exotic layouts) to compile Python functions into XLA programs. The Mesh is supplied via in_axes and out_axes arguments, causing XLA to generate a single executable that spans the entire TPU pod.

The compilation happens in [levanter/trainer.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/trainer.py), where the trainer coordinates the transition from eager Python to compiled SPMD code:

import jax
import optax

def train_step(state, batch):
    def loss_fn(params):
        logits = state.model.apply(params, batch["inputs"])
        loss = optax.softmax_cross_entropy_with_integer_labels(
            logits, batch["labels"]
        ).mean()
        return loss

    grads = jax.grad(loss_fn)(state.params)
    grads = jax.lax.pmean(grads, axis_name="devices")  # sync across pod

    new_params = optax.apply_updates(state.params, grads)
    return state.replace(params=new_params)

# compiled once for the whole pod

pjit_train_step = jax.pjit(
    train_step,
    in_shardings=(None, None),   # let Levanter infer from mesh

    out_shardings=None,
    donate_argnums=(0,),
    backend="tpu",
)

Multi-Host Synchronization and State Management

Training across multiple TPU hosts requires careful coordination to ensure model parameters and optimizer state remain consistent.

Cross-Replica Communication with psum and pmean

For multi-host pods, levanter.trainer uses jax.lax.psum and jax.lax.pmean to aggregate gradients and metrics across the entire device mesh. These collective operations ensure that every replica maintains identical parameter values after each update step.

Barrier Coordination in trainer_state.py

A small "coordinator" process runs a barrier on each host to guarantee that all replicas finish XLA compilation before executing the first training step. This synchronization logic resides in [levanter/trainer_state.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/trainer_state.py), preventing race conditions during the transition from compilation to execution.

TPU-Specific Optimizations and Kernel Selection

Levanter ships with optimized kernels that leverage TPU hardware features, falling back to reference implementations when unavailable.

Splash Attention and Pallas Kernels

The library optionally enables Splash Attention and Pallas kernels when levanter.utils.cloud_utils.is_tpu detects a TPU backend. According to [levanter/layers/attention.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/layers/attention.py), these kernels provide high-performance attention computations optimized for TPU memory hierarchies, with automatic graceful degradation to standard JAX implementations on other hardware.

Resource Management and Cleanup

Levanter includes utilities to manage cloud costs by handling VM lifecycle automatically.

Automatic VM Shutdown

After training completes, the Trainer optionally shuts down the TPU VM via self.shutdown_at_exit. This feature uses the Cloud TPU metadata service to terminate resources, ensuring predictable cloud costs for batch training jobs. The shutdown logic is implemented in [levanter/trainer.py](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/trainer.py).

End-to-End Training Example

To launch a distributed training job, users call entry points like train_lm.py with TPU-specific parameters:


# example: launch a tiny LM on a full-pod TPU

from levanter.main.train_lm import main as train_lm

if __name__ == "__main__":
    train_lm(
        model_config="tiny_gpt2",
        dataset="tiny_corpus",
        batch_size=1024,          # per-device batch size

        epochs=10,
        tpu=True,                 # request TPU resources

        mesh_shape="2x4x8",       # optional override of TPU topology

    )

This single command triggers the full compilation pipeline: topology detection, mesh construction, pjit compilation, and distributed execution across the entire TPU pod.

Summary

  • Automatic topology detection: hardware_topology.py reads TPU slice metadata to build the device grid.
  • Global mesh construction: distributed.py creates a jax.sharding.Mesh visible to all training components.
  • SPMD compilation: trainer.py uses jax.pjit with mesh sharding to generate pod-wide XLA executables.
  • Cross-host synchronization: trainer_state.py coordinates barriers and collective operations (pmean, psum) to maintain consistency.
  • Hardware optimization: attention.py automatically enables TPU-specific kernels with fallback support.
  • Resource cleanup: Optional automatic VM shutdown prevents unnecessary cloud costs after training completion.

Frequently Asked Questions

How does Levanter detect TPU topology automatically?

Levanter queries the Cloud TPU metadata service through levanter.utils.hardware_topology to read the physical slice layout (e.g., 2x4x8). This module parses the AxBxC or SxAxBxC topology strings automatically, eliminating the need for manual mesh configuration in standard deployments.

What is the role of jax.pjit in Levanter's distributed strategy?

jax.pjit compiles Python training functions into single-program multi-data (SPMD) XLA executables. In Levanter, pjit uses the global Mesh constructed from hardware detection to shard data and parameters across TPU cores, enabling the same Python code to run efficiently on thousands of devices simultaneously.

How does Levanter handle synchronization across multiple TPU hosts?

The library uses jax.lax.pmean and jax.lax.psum for gradient and metric aggregation across the device mesh. Additionally, trainer_state.py implements a barrier coordinator that ensures all hosts complete XLA compilation before the first training step begins, preventing desynchronization errors.

Can Levanter training jobs fall back if TPU-specific kernels are unavailable?

Yes. The levanter.layers.attention module checks levanter.utils.cloud_utils.is_tpu to conditionally enable Splash Attention and Pallas kernels. If the TPU backend is not detected or specific kernels are incompatible, the code automatically falls back to reference JAX implementations, ensuring portability across hardware platforms.

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 →