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/haliaxfor high-performance tensor operations andlib/marin/src/marin/processing/tokenize/for tokenization pipelines. -
Training Core: Implements distributed JAX training via
lib/levanter, with optimizer logic located inlib/levanter/optimand 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 bylib/haliaxutilities. -
Orchestration: Manages experiment steps and resource allocation through
lib/irisfor cluster job orchestration andlib/marin/src/marin/execution/step_runner.pyfor step resolution. -
Evaluation & Reporting: Runs benchmarks and generates model cards via modules in
lib/marin/src/marin/evaluation/and documentation indocs/reports/. -
Infrastructure: Deploys pipelines on TPU/GPU clusters using
infra/(Pulumi IaC) and profiles performance viainfra/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:
-
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.
-
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.
-
Run the graph using a
StepRunnerthat 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:
-
README.md– High-level project description and entry points for the repository. -
lib/levanter/README.md– Documentation for the JAX-based training framework, including distributed training strategies. -
lib/haliax/README.md– Overview of the tensor library used for model sharding and high-performance computation. -
experiments/tutorials/train_tiny_model.py– Complete reference implementation demonstrating the full training pipeline. -
lib/marin/src/marin/execution/step_runner.py– Core executor that materializes step graphs and manages resource allocation. -
lib/marin/src/marin/experiment/data.py– Lazy dataset definition utilities and tokenization interfaces. -
lib/marin/src/marin/experiment/train.py– High-level training step builder that configures optimizers and model architectures. -
infra/pulumi/README.md– Infrastructure-as-code scripts for cluster provisioning and deployment.
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
StepRunnerclass inlib/marin/src/marin/execution/step_runner.pyautomatically 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →