Building with AI

Reproducing OLMo 3 7B on TPUs reveals hidden data loader bugs

A Google case study reproduces AI2's OLMo 3 model on TPUs, matching performance metrics while uncovering critical data sharding errors that masked memorization as improvement.

A glass cube with glowing circuits representing a TPU model next to paper notes.
Image: Google Developers Blog, licensed CC BY 4.0

In a post on the Google Developers Blog in September 2026, engineers detailed their successful reproduction of the Allen Institute for AI’s OLMo 3 7B language model using MaxText on Google Cloud TPUs. The team matched the original training curves and downstream benchmarks, proving that a PyTorch-based GPU recipe could be faithfully executed in a JAX-based TPU environment.

What happened

The engineering team set out to validate the capabilities of MaxText, a high-performance framework for training large language models on TPUs. They chose OLMo 3 because it is one of the few frontier-class models with fully open weights, data, and training recipes. This transparency allowed for a rigorous comparison against an independent reference run, rather than relying solely on loss curves. The reproduction covered stage-1 pre-training and stage-2 mid-training annealing, utilizing AI2’s step-0 PyTorch checkpoints converted to the Orbax format for JAX.

The results showed that MaxText could track AI2’s published loss curve over the full 5.93 trillion-token budget. However, the process was not without complications. During the later stages of training, the MaxText run appeared to outperform the reference, with training loss dropping significantly below AI2’s levels. Further investigation revealed this was not a genuine improvement in model quality but a artifact of a data loading bug. The team also demonstrated the system's resilience by resizing the training job mid-flight after losing hardware capacity and switching TPU generations between training stages without altering the core recipe.

How it works

To ensure fidelity, the team ported OLMo 3’s specific architectural choices, including reordered-norm blocks, QK-norm, and a mixed sliding-window and global attention pattern, into MaxText. They verified the conversion by checking logit parity, ensuring the converted model matched the HuggingFace reference within the expected noise floor for different frameworks. The data pipeline was rebuilt using Grain, implementing random-access reads, global index shuffling, and n-gram repetition filtering to mirror the original setup exactly.

Figure from the original article: Reproducing OLMo 3 7B on TPUs reveals hidden data loader bugs
Figure from the original article · Google Developers Blog · CC BY 4.0

Performance optimization played a crucial role in making the reproduction feasible. The team achieved 44.5% model flops utilization (MFU) on Ironwood TPUs by offloading collectives to the SparseCore and tuning rematerialization strategies. They also explored co-design opportunities, such as reshaping attention heads to better utilize the TPU’s matrix multiplication units, which yielded faster training speeds without sacrificing quality. These engineering adjustments allowed the system to maintain throughput even when hardware resources changed dynamically during the multi-week training run.

Key details

  • Metric matching: The reproduced model matched AI2’s held-out evaluation metrics, including C4 loss and eight downstream tasks, with accuracy differences never exceeding ±0.005.
  • Data bug discovery: A double-sharding error in the Grain data loader caused repeated data instances, artificially lowering training loss by up to 0.25 points without improving generalization.
  • Hardware resilience: The training job successfully resumed on a cluster slice one-quarter the original size after a capacity loss, maintaining per-device throughput within 1%.
  • Framework conversion: PyTorch weights were converted to Orbax with a KL divergence of approximately 1.5e-3, confirming the models were identical at initialization.
  • Simplified recipe: The team used a single cosine learning rate schedule instead of AI2’s stitched two-curve approach, yet still achieved comparable results.

Why it matters

For developers building large-scale AI systems, this case study highlights the danger of relying exclusively on training loss as a proxy for model quality. The incident where MaxText appeared to beat the reference due to data memorization serves as a stark reminder that lower loss does not always mean better performance. It underscores the necessity of rigorous held-out evaluations and independent verification steps to catch subtle bugs in data pipelines that can otherwise go unnoticed for weeks.

Figure from the original article: Reproducing OLMo 3 7B on TPUs reveals hidden data loader bugs
Figure from the original article · Google Developers Blog · CC BY 4.0

Furthermore, the successful reproduction demonstrates the maturity of JAX and TPU ecosystems for handling complex, non-standard transformer architectures. It provides a blueprint for teams looking to migrate from PyTorch/GPU stacks to JAX/TPU environments, showing that faithful reproduction is possible even with significant architectural differences. The ability to resize jobs and switch hardware generations mid-training also offers practical insights for managing cost and reliability in long-running cloud training jobs.

What you can do

  • Implement regular held-out evaluations alongside training loss monitoring to detect memorization or data leakage early.
  • Verify data loader sharding logic with unit tests that simulate multi-worker environments to prevent silent data duplication.
  • Use logit parity checks when converting models between frameworks to ensure initialization consistency before starting training.
  • Design training scripts to handle dynamic resource changes, allowing for seamless resumption on different hardware slices if preemptions occur.
  • Compare downstream task accuracy at multiple checkpoints, not just at the end of training, to ensure consistent capability growth.
  • Review data pipeline configurations for redundant sharding operations, especially when using libraries like Grain that may handle sharding internally.

Tools from the Bytechap store

$89

DocBento

Self-hosted document management that reads every scan and answers with page citations.

Live demo

Keep reading

All stories