Google's TPU Team Reproduces Ai2's Olmo 3 7B Training Run in MaxText
In short: A Google Cloud TPU engineering team worked with Ai2 to rebuild Olmo 3 7B's pre-training from scratch using MaxText, Google's JAX/XLA training framework, running on TPUs instead of Ai2's original PyTorch/GPU setup. They matched Ai2's published results not just on the training loss curve but on four independent held-out evaluation surfaces across the full ~5.93-trillion-token, 1.41-million-step stage-1 run plus the stage-2 annealing phase. Along the way they ported Olmo 3's unusual architecture (reordered-norm blocks, QK-norm, 3:1 sliding/global attention ratio) into MaxText and caught a data-loader bug that had been quietly inflating apparent performance through memorization rather than genuine learning.
This summary was generated automatically by AI from Google Developers's publication. It is our own text, not a copy of the original — facts, figures and quotes belong to the source, linked above and below.
What changed?
- 1Reproduced Olmo 3 7B stage-1 pre-training (~5.93T tokens, 1.41M steps) and stage-2 mid-training anneal in MaxText on TPUs, matching Ai2's PyTorch/GPU reference
- 2Verified the match on four held-out metrics: C4 held-out loss, an 8-task lm-eval-harness suite (MMLU, HellaSwag, ARC, OpenBookQA, PIQA, BoolQ, WinoGrande), multi-domain perplexity, and token-level KL divergence — not just training loss
- 3Ported Olmo 3's non-standard architecture (reordered-norm, QK-norm, 3:1 sliding/global attention) to JAX, verified via a logit-parity check (KL ≈ 1.5e-3, 98.75% top-1 token agreement at 8192-token context)
- 4Found and fixed a Grain data-loader double-sharding bug that caused a Poisson-style resampling, making ~26% of the corpus seen twice or more — this deflated training loss through memorization while held-out metrics stayed flat
- 5Demonstrated checkpoint-and-resume reliability (Δ=0.000 across a controlled resume, and after an actual host failure the run re-trained 127 steps with Δ=0.000)
- 6Resized the training job mid-run after losing 75% of TPU capacity, resuming on a quarter-size slice with no recipe change and <1% per-device throughput loss
- 7Switched TPU generations mid-recipe (Ironwood to v5p) for stage 2 using the identical launcher, sustaining 57.4% MFU
- 8Achieved 44.5% Model FLOPs Utilization on Ironwood via SparseCore collective offload, remat tuning, and sharding optimization
- 9A side ablation reshaping attention heads (32×128 to 16×256) ran 12.4% faster at identical quality by better utilizing Ironwood's 256×256 matrix unit
Why it matters
This is a rigorous engineering validation that Google's MaxText/TPU training stack can reproduce a known-good open model recipe end-to-end, with the same generalization behavior as the original — not just a similar-looking loss curve. It also shows a concrete failure mode (a data pipeline sharding bug masquerading as a training improvement) that any team running large-scale LLM pre-training or fine-tuning should guard against by evaluating held-out metrics, not training loss alone.
Sources
- Google DevelopersOfficialPrimary sourceOriginal article →„Reproducing Olmo 3 7B Pre-training in MaxText: case study of large scale training on TPUs“24 Sept 2026, 03:00
- Published by source
- 24 Sept 2026, 03:00
- Found by our system
- 2 Oct 2026, 19:51
- Summary generated
- 2 Oct 2026, 22:56
This article was written by AI from the original source. Facts, numbers and prices come from the source; missing values are marked “Not specified”. Legal notice, copyright and privacy