
7/16/2025 · Srikanth Kilaru, David Hall
What this post added
This post details how the Marin project leverages JAX and its associated frameworks (Levanter, Haliax) to build and train large-scale, reproducible foundation models. Key technical contributions include: 1. Encapsulating training steps into single `@jax.jit`-decorated functions for fused operations and performance optimization. 2. Utilizing `jax.value_and_grad` for efficient loss and gradient computation, and `gradient checkpointing` for memory savings. 3. Employing Pallas-based Splash Attention for optimized Dot Product Attention. 4. Leveraging `@jax.jit` for SPMD parallelization and automated sharding/communication. 5. Introducing Haliax (named tensors) within Levanter for simplified management of complex sharding strategies like FSDP and Tensor Parallelism. 6. Integrating Google Cloud TPU Multislice and Ray for resilient, cost-effective compute cluster management using preemptible instances. 7. Emphasizing JAX's reproducibility guarantees (deterministic PRNGs) and using Tensorstore for deterministic data loading, enabling bit-for-bit reproducibility across training migrations and preemptions. The post also describes the Llama-style transformer architecture of Marin-8B and its adaptive training process ('Tootsie' process) across different TPU configurations.