How to Build Foundation Models with Marin: A Complete Guide to the Open-Source Training Stack

Marin is an open-source research framework that uses lazy step handles, JAX-based distributed training, and immutable configurations to build reproducible foundation models at any scale.

Marin is an open-source research program that provides a full stack for building, training, and evaluating large-scale foundation models. The repository at marin-community/marin separates concerns across data processing, model architecture, and infrastructure orchestration using a layered architecture. When you build foundation models with Marin, you define computational graphs where each step executes only after its dependencies are satisfied, enabling reproducible experiments that scale from CPU prototypes to TPU clusters.

Understanding Marin's Layered Architecture for Foundation Models

Marin organizes its foundation model pipeline into six distinct layers, each implemented in specific packages and modules:

  • Data & Tokenization: Curates and tokenizes raw data using lib/haliax for high-performance tensor operations and lib/marin/src/marin/processing/tokenize/ for tokenization pipelines.

  • Training Core: Implements distributed JAX training via lib/levanter, with optimizer logic located in lib/levanter/optim and mixed-precision support for large-scale training.

  • Model & Architecture: Defines model families including MoE and decoder-only LLMs in lib/marin/src/marin/modeling/, with sharding strategies provided by lib/haliax utilities.

  • Orchestration: Manages experiment steps and resource allocation through lib/iris for cluster job orchestration and lib/marin/src/marin/execution/step_runner.py for step resolution.

  • Evaluation & Reporting: Runs benchmarks and generates model cards via modules in lib/marin/src/marin/evaluation/ and documentation in docs/reports/.

  • Infrastructure: Deploys pipelines on TPU/GPU clusters using infra/ (Pulumi IaC) and profiles performance via infra/xprof/.

All components couple through lazy step handles, allowing you to declare a graph of tasks where each step materializes only after its dependencies complete. This design mirrors a Makefile but provides Python-level composability and automatic resource scheduling.

The Lazy Execution Model: How Marin Handles Dependencies

Marin's execution model treats foundation model training as a directed acyclic graph of immutable steps. Each step is versioned using name plus version parameters, and configurations are immutable data classes stored in a central metadata store.

The typical workflow follows three stages:

  1. Define a tokenized dataset as a lazy handle pointing to a tokenized version of raw data (e.g., TinyStories). No data downloads until downstream steps request it.

  2. Create a training step that specifies model architecture, optimizer, batch size, sequence length, and training steps. This step automatically pulls in the tokenized dataset handle as a dependency.

  3. Run the graph using a StepRunner that resolves dependencies, provisions required resources (CPU, GPU, or TPU), and executes steps in topological order.

This approach ensures that checkpoints, logs, and metrics remain reproducible across runs, with all artifacts stored in the metadata store.

Step-by-Step Tutorial: Training Your First Foundation Model

The following example demonstrates how to build a tiny language model on the public TinyStories dataset using Marin's core concepts.

Step 1: Define a Tokenized Dataset

First, create a lazy handle to the dataset. The actual download and tokenization occur only when the training step reads from it.

from marin.experiment.data import tokenized
from experiments.marin_tokenizer import marin_tokenizer

# Lazy tokenization – no data is fetched yet.

tinystories_tokenized = tokenized(
    name="tokenized/tinystories",
    source="roneneldan/TinyStories",
    tokenizer=marin_tokenizer,
    sample_count=1000,          # small sample for a quick tutorial

)

Step 2: Configure the Training Step

Next, define the training configuration using the train_lm function from lib/marin/src/marin/experiment/train.py. This step consumes the tokenized dataset handle and specifies hardware resources via ResourceConfig.

from fray.cluster import ResourceConfig
from levanter.optim import AdamConfig
from marin.experiment.train import train_lm
from experiments.llama import llama_nano

# Training step – depends on the tokenized dataset above.

nano_tinystories_model = train_lm(
    name="checkpoints/marin-nano-tinystories",
    version="v1",
    model=llama_nano,
    optimizer=AdamConfig(learning_rate=6e-4, weight_decay=0.1),
    datasets={tinystories_tokenized: 1.0},
    batch_size=4,
    seq_len=2048,
    num_train_steps=100,
    resources=ResourceConfig.with_cpu(),
)

Step 3: Execute the Computation Graph

Finally, use the StepRunner from lib/marin/src/marin/execution/step_runner.py to resolve the dependency graph and execute the training.

from marin.execution.lazy import lower
from marin.execution.step_runner import StepRunner

# Execute the graph.

if __name__ == "__main__":
    StepRunner().run([lower(nano_tinystories_model)])

The lower() function transforms the lazy step into an executable format, while StepRunner().run() provisions the CPU resources and processes the steps in dependency order.

Scaling from Nano Models to Production MoE Architectures

The same pattern scales to massive mixture-of-experts (MoE) models. To scale up, replace llama_nano with an MoE architecture from lib/levanter or lib/haliax, increase batch_size and seq_len to handle longer contexts, and switch to TPU resources via ResourceConfig.with_tpu().

The sharding utilities in lib/haliax automatically handle model parallelism across devices, while lib/levanter manages the distributed JAX training loop. For production deployments, the infra/ directory contains Pulumi-based infrastructure-as-code scripts for provisioning TPU pods or GPU clusters.

Key Source Files for Building Foundation Models

When working with Marin, you will interact with these critical modules:

Summary

  • Marin provides a six-layer architecture separating data processing, training, modeling, orchestration, evaluation, and infrastructure for foundation model development.

  • Lazy step handles enable declarative workflow graphs where steps execute only after dependencies complete, ensuring efficient resource utilization.

  • Immutable configurations and versioned steps in lib/marin/src/marin/execution/ guarantee reproducibility across experiments.

  • Levanter (lib/levanter) and Haliax (lib/haliax) provide the JAX-based training core and tensor operations required for distributed foundation model training.

  • The StepRunner class in lib/marin/src/marin/execution/step_runner.py automatically provisions CPU, GPU, or TPU resources and executes steps in topological order.

  • Workflows scale seamlessly from nano models on CPUs to production MoE architectures on TPU clusters by changing resource configurations and model imports.

Frequently Asked Questions

What hardware accelerators does Marin support for training foundation models?

Marin supports CPU, GPU, and TPU clusters through the ResourceConfig class in fray.cluster. You specify hardware requirements when defining training steps—for example, ResourceConfig.with_cpu() for local development or ResourceConfig.with_tpu() for large-scale distributed training. The infra/ directory contains Pulumi scripts for provisioning these resources on cloud platforms.

How does Marin ensure reproducibility in foundation model experiments?

Marin ensures reproducibility through three mechanisms: every step is versioned using name and version parameters, configurations are immutable data classes (such as AdamConfig), and all artifacts including checkpoints and metrics are stored in a central metadata store. The lazy execution model guarantees that data processing and training steps execute deterministically based on their dependency graph.

Can I use custom model architectures not defined in the Marin repository?

Yes. You can import custom architectures by defining them in modules like experiments/ and passing them to the train_lm function. Marin expects models to be compatible with JAX and Haliax sharding utilities, but the framework does not restrict you to predefined architectures in lib/marin/src/marin/modeling/. As long as your architecture integrates with the Levanter training loop, you can use it within the Marin execution graph.

What is the role of Haliax in Marin's foundation model stack?

Haliax (lib/haliax) serves as Marin's high-performance tensor library, providing the computational primitives and sharding utilities required for distributed training. It handles model parallelism and device mesh strategies, enabling efficient scaling of large foundation models across multiple accelerators. The library is used throughout the data processing and model architecture layers to manage tensor operations and memory layout.

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 →