OLMo 3 由人工智能艾伦研究所(AI2)开发,是一个最先进的、完全开源的语言模型,采用现代架构和多阶段训练配方进行训练。为了评估 MaxText 在 Google Cloud TPUs 上的能力,我们的团队着手从头复现 AI2 的 OLMo 3 7B 模型。我们选择 OLMo 3 是因为它结合了三个很少同时出现的特性:它是一个强大的、现代的 7B 模型,以真实的生产规模进行训练;AI2 公开了几乎完整的模型流程,包括数据、代码、配置、检查点、日志和评估结果;最后,它为我们提供了一个独立的 PyTorch 和 GPU 参考基准,我们可以据此测试 MaxText 和 TPUs。
我们在 Google Cloud TPUs 上的 MaxText 中复现了 AI2 的 OLMo 3 7B,包括阶段 1 预训练和阶段 2 中期训练退火,并在保留指标上证明了匹配性,而不仅仅是损失曲线:
主要亮点如下,每一点将在后文中详细阐述:
PyTorch → JAX 模型转换。OLMo 3 的架构(重排序归一化块、QK 归一化、3:1 滑动/全局注意力)被移植到 MaxText 中,并通过 logits 一致性检查进行了验证:转换后的第 0 步检查点与 HuggingFace 参考在 KL 散度 ≈ 1.5e-3 处匹配,这处于“同一模型,不同框架”的噪声基线水平;在 bfloat16 全 8192 token 上下文下,两者在 top-1 token 上的一致性达到 98.75%。
能够发现真实错误的验证过程。保留评估发现了一个数据加载器 bug,该 bug 使 MaxText 看起来优于参考模型;这种优势实际上是记忆效应所致。
多周运行的可靠性。检查点与恢复重放完全复现了运行过程:受控的 A/B 测试显示,在恢复后的每一步,Δ = 0.000;当主机故障导致阶段 2 运行中途终止时,恢复后的运行在记录的损失和困惑度上以 Δ = 0.000 重新训练了 127 步。
运行时调整训练作业规模。在第约 1.05M 步时,我们失去了四分之三的容量;运行在仅为四分之一大小的切片上恢复,无需更改配方(相同的脚本 run_olmo3_7b_stage1.sh 按设备批量大小缩放以保持 GBS 恒定),每个设备的吞吐量保持在 1% 以内(≈100% 强扩展性,双向测量)。
在配方中途更换 TPU 代际。阶段 2 将相同的启动器指向 v5p 而非 Ironwood,仅更改了设备类型,并保持了 57.4% 的 MFU。
物有所值的性能优化。通过 SparseCore 集体卸载、remat 调优和最佳分片,在 Ironwood 上以 7B 规模实现了 44.5% 的 MFU,大致收回了三分之一的计算预算。
为 TPU 进行的协同设计:在相同质量下更快。将注意力机制从 32 个头 × 头维度 128 重塑为 16 × 256,在参数和 FLOPs 相同的情况下,运行速度提高了 12.4%,因为头维度 256 充分利用了 Ironwood 的 256×256 MXU,且其损失曲线在前 120B token(30k 步)内与原始曲线匹配。这是一项侧向消融实验;复现过程保留了原始架构。
从 AI2 的第 0 步 PyTorch 权重和相同的核心配方开始,MaxText 的运行在整个约 5.93T token / 1.41M 步的预算中追踪了 AI2 发布的损失曲线,并在阶段 1 结束时与其重合。我们甚至简化了两个配方细节(AI2 拼接了两个余弦学习率调度,而我们使用单一调度;以及公开发布的数据混合;见下方的配方),但匹配性依然成立。本文的其余部分将介绍每一项是如何构建、测量以及在其中一个具有启发性的案例中几乎被伪造的。
为什么要复现 OLMo 3?
OLMo 3 是少数几个真正开源的前沿级语言模型之一:开放权重、开放数据,以及完全指定的训练配方,并在 Weights & Biases 上提供了公开的参考运行。独立复现该运行,并在保留指标而非仅仅是损失曲线上进行匹配,是强有力的证据,证明 MaxText 堆栈(优化器、损失函数、数据管道、数值计算)是忠实的,而不仅仅是“看起来像是在训练”。
MaxText 是一个专为 TPU 构建的 JAX/XLLM 训练框架。我们要回答的问题是:能否在 JAX-on-TPU 上忠实地复现 PyTorch-on-GPU 的训练配方,在关键指标上实现匹配(而非逐位精确匹配),以及如何证明这一点?
OLMo 3 的训练配方是一个三阶段课程:通用预训练、中期训练(退火)和长上下文适应。本文涵盖第 1 阶段(约 5.9T token 的预训练运行)和第 2 阶段(中期训练),两者均端到端训练,并与 AI2 的参考结果进行匹配。第 3 阶段及后训练(通过 Tunix 进行的 SFT/RL)是我们已编写但尚未运行的配方。
OLMo-3 预训练课程:本文复现了第 1 阶段(Ironwood)和第 2 阶段退火(v5p);第 3 阶段长上下文(seq 65k, YaRN)及后训练(通过 Tunix 进行的 SFT 然后 GRPO)将在后续介绍。
OLMo-3 课程。本文复现了第 1 和第 2 阶段;第 3 阶段及后训练将在后续介绍。
配方
OLMo 3 7B 是一个具有 32 层、4096 维的密集 Transformer,包含一些非标准选择:“重排序归一化”块、QK 归一化以及 3:1 的滑动窗口与全局注意力混合。MaxText 配置(olmo3-7b-pt.yml,用于第 1 和第 2 阶段)与其完全匹配:
训练配方镜像了 OLMo-core 的 pretrain-1.py;为了使曲线对齐必须匹配的旋钮包括:
我们从 AI2 的 step-0 PyTorch 检查点开始训练,并将其转换为 Orbax 格式,因此 MaxText 从与参考模型完全相同的权重开始。转换本身构成了第一个检查点:在转换后的权重上进行的前向传播与 HuggingFace 参考结果匹配,KL 散度约为 1.5e-3,前 10 个 token 的重叠率为 9/10,这体现了“相同模型,不同框架”的噪声底限。
这种匹配是否依赖于继承 AI2 的初始化?显然并非如此。作为独立检查,我们还使用 MaxText 自身的随机初始化训练了约 50k 步(占总步数的 3.5%);其训练损失紧密跟踪了 AI2 发布的曲线,略低于该曲线。这是一次训练损失的抽样检查,而非完整复现,但它表明匹配并不取决于是否从 AI2 的权重开始。
数据管道与 OLMo-core 完全镜像:对所有文档进行分词并连接(EOS 之间),切片为不重叠的 8192 token 实例,使用固定种子对索引进行全局洗牌,并应用 n-gram 重复过滤器,屏蔽包含超过 32 个重复 n-gram 的实例。MaxText 的 dataset_type=olmo_grain(基于 Grain 构建)实现了这一点。
与 AI2 运行相比的两个有意差异。(1)学习率(LR)调度:AI2 最初计划使用约 5T token,并在中途将运行延长至最终目标约 5.93T,因此其 LR 轨迹拼接了两个余弦曲线(可见于其公开的 WandB 运行);我们在全目标范围内运行了单个余弦曲线。(2)数据:我们在公开发布的 OLMo-3 混合数据集上进行训练,该数据集省略了 AI2 内部运行中看到的一小部分数据(不到 token 预算的 0.5%,主要是未包含在发布文件列表中的 s2pdf 分片)。这两者都是我们主动选择的简化措施,而非意外结果,MaxText 在所有保留表面上仍然匹配。这也是为什么我们说“复现精度在运行间噪声范围内”,而非逐位精确匹配(见 §G 中的 KL 分析)。
我们需要构建的内容
当我们开始时,OLMo 3 并不在 MaxText 中;此次复现添加并合并了以下内容。“本文末尾的‘自行复现’”指的是配置,而非代码。
模型本身:重排序归一化块、QK 归一化以及 3:1 的滑动/全局注意力模式(#3004, #3112)。
与 OLMo-core 语义完全匹配的跳过步优化器,精确到贝塞尔校正的运行标准差(在 128 步窗口内超过 6σ 时跳过)(#3490)。
z-loss(#3211)和每参数权重衰减掩码,以便排除嵌入层(#3280)。
olmo_grain 数据管道(#3749):预分片分片的随机访问读取、带有指纹保护以防止重启时静默数据交换的带种子全局索引洗牌、n-gram 重复过滤器,以及(从第 2 阶段开始)Grain 迭代器状态检查点。
HF↔Orbax 检查点转换配合 logit-parity 校验,这是本文所有“框架噪声基底”数值的背后工具(#3112,#3832)。
阶段 1 和 2 启动器(环境驱动的运行脚本 + 带有 submit / monitor / resume_until_done 功能的 XPK 包装器)(#3886)。
TensorBoard 一致性标签(optim/step_skipped, perf/total_tokens),使得 AI2 的 W&B 仪表板上的每个指标都有一个 MaxText 对应项可供比较。
它收敛了吗?
头条新闻是一个叠加层:MaxText 的阶段 1 lm_loss 与 AI2 发布的 WandB 曲线,步骤对齐并分箱为 2k 步均值。在约 80 万步内,两条曲线保持在 ±0.012 以内;从约 90 万步开始,MaxText 低于 AI2 且从未再交叉回来,这是下一节剖析的数据 bug 的第一个迹象。
MaxText 与 AI2 阶段 1 损失在 141 万步上的对比,下方面板显示经过 18k 步平滑处理的差距
顶部:在此尺度下,两条曲线直到尾部都难以区分。底部:差距在约 80 万步内保持在 ±0.012 以内,从约 90 万步开始因数据 bug 重复累积而向负方向倾斜,并在超过 125 万步后急剧下降,在原始 2k 步分箱中达到 −0.22,随后才应用两个面板绘图所用的 18k 步平滑处理。这些数据均未触及保留集损失或下游准确率,如下文各节所示。(每地标表:附录 A;完整曲线已提交为 olmo_stage1_loss_curve.tsv。)
但仅凭损失曲线是薄弱的证据:两次运行可能在训练损失上匹配,但在你真正关心的所有方面产生分歧。因此,我们在跨越 91.5 万步的六个步骤地标上,在四个独立的表面上验证了收敛性:
保留集 C4 lm_loss :对 1600 万个 token 的保留集 C4-en 进行仅前向传播评估,两次运行使用相同的批次。
8 任务 lm-eval-harness 套件:MMLU, HellaSwag, ARC-easy/challenge, OpenBookQA, PIQA, BoolQ, WinoGrande。
多域保留集困惑度:在网页、新闻、百科全书和混合领域进行的 Paloma 风格扫描。
Token 级 KL:相同输入下的下一个 token 分布距离。
测量方法。所有评估均在步骤对齐的检查点对上运行:实时 MaxText 运行与同一步骤的 AI2 检查点,即其公共 HuggingFace 版本 allenai/Olmo-3-1025-7B@stage1-step{N} 转换为 Orbax(该转换重现了 PyTorch 参考值,KL ≤ 1.8e-3,即框架噪声基底)。lm-eval 使用标准的 lm-eval-harness(5-shot MMLU,其他默认);σ 是每个任务 harness 的标准误,以平方和形式组合用于计算差异。
在阶段 1 结束时,所有表面都表明这两种配方是可互换的:
第四个表面,token 级 KL,是唯一一个数值不为极小的指标:在相同输入上的平均值为 0.389 nats,约为框架噪声基底的 200 倍。这对于同一配方的两次独立运行(相同的总体技能,不同的概率质量分配)是预期的,这也是我们说“运行间噪声”而非逐位一致的原因;详细分解见 §G 。
下游准确率差距在所有六个地标上从未超过 ±0.005 的宏观值,符号翻转了四次,这正是两个仅在随机数生成器和数值计算上存在差异的忠实运行所预期的随机游走(每地标表见附录 D):
8 任务宏观准确率,MaxText 与 AI2 在六个地标上的对比,带有每地标的差异
在每个地标上,下游能力是可互换的。与训练损失不同,准确率从未单调发散;差异在 ±0.005 内部进行随机游走,最终结束于 +0.0002。
看似胜利实则 bug
这里变得有趣了。从约 90 万步开始,MaxText 的训练损失低于 AI2 且从未再交叉回来;超过 125 万步后,它明显低于 AI2,平均低 −0.06,在某些几百步的区间内低至 −0.25。如果只看训练损失叠加层,你会得出结论认为 MaxText 已经领先。
并没有。在边界检查点处的保留 C4 损失持平(在 100 万步时为 Δ −0.004,在第一阶段结束时为 +0.003),而在 135 万步时的下游准确率略微偏向 AI2 。训练损失正在下降,但泛化能力没有变化。这是记忆化的典型特征:模型看到了某些序列超过一次,并在重复出现时获得了较低的损失。
上图:MaxText 的训练损失降至 1.63,而 AI2 保持平稳。下图:保留的 C4 差异在整个过程中保持在零附近。
一张图概括全部故事。上图:在 124 万至 141 万的窗口期内,MaxText 的训练损失在 2000 步均值上反复 plunge 至 1.63(Δ −0.22;在百步区间内为 −0.25),而 AI2 的损失保持平稳。这看起来像是一场失控的胜利。下图:训练损失差异(蓝色)向负值漂移,但保留的 C4 损失(绿色菱形)从未离开 ±0.02 的范围。训练损失下降了;泛化能力没有。
原因是 Grain 数据加载器中存在一个双重分片错误。MaxText 的 OLMo 加载器将 ShardOptions(shard_index, shard_count) 传递给 Grain DataLoader,而索引采样器已经在内部进行分片。Grain 的 shard_options 不仅仅记录元数据;它还会重新步长化采样器的索引流。当 shard_count=32 时,数据游标前进速度快了 32 倍,因此第一阶段不再是一个干净的 epoch,而是变成了泊松分布 (≈1) 的有放回重采样:大约 37% 的语料库从未被看到,37% 被看到一次,26% 被看到两次或更多。令牌预算保持不变(约 5.9T 真实令牌),这就是为什么损失在整体上仍然与 AI2 保持一致的原因,但重复出现的实例在其发生处恰好压低了训练损失。
由此得出了两条教训:
训练损失不是收敛的证明。我们没有发布虚假的“MaxText 优于参考”声明的唯一原因是,我们承诺在每个关键节点进行保留评估。记忆化导致的低谷在保留的 C4 和所有 8 个下游任务上是不可见的。
诚实的重现需要预留处理错误的预算。修复方法(grain.sharding.NoSharding(),让采样器拥有所有分片逻辑)只有一行代码。找到它需要一个 A/B 测试框架、一个能在 shard_count>1 时重现差异的单元测试,以及一次硬件重运行以进行验证。
在验证修复方案的过程中,我们发现了一个第二独立的错误:恢复步骤检测中的 off-by-one 错误。检查点目录编号为 N,但训练循环在完成第 N 次迭代后写入目录 N,因此模型恢复到第 N+1 步,而数据加载器在第 N 个批次处恢复,导致重新训练一个批次,然后永久落后一步。在两个修复方案都应用后,检查点并恢复的运行精确复现了不间断运行的结果:在所有 99 个步骤中,记录的损失差异 Δ = 0.000。(该 A/B 测试运行的是单工作器数据加载;多工作器情况出现在第二阶段,我们已在其中关闭了它;见第二阶段。)这两个错误都有回归测试,旧代码会导致测试失败,而修复后的代码会使测试通过。
我们让正在进行的阶段 1 运行按原样完成:它已经完成了 85%,修复方案无法重新打乱已经读取的数据,而且重新启动将浪费约 120 万步的计算资源。上述验证表明,该错误对可观察的准确率没有造成任何影响;此修复方案用于未来的运行。
Ironwood 上的性能和扩展性
复现数学计算只是工作的一半;另一半是让运行速度快,并在集群环境发生变化时保持高速。在为期数周、完成 140 万步运行的过程中,任务被抢占、重新调度并调整大小多次,而软件栈必须吸收所有这些变化,同时不触碰配方(配置)。
榨取 MFU
在 Ironwood 上,模型大小为 7B,每设备批次为 4,我们在基础架构(“变体 D”,这是我们在附录 J 的消融实验中使用的标签)下达到了 44.5% 的 MFU(每台设备 510–513 TFLOP/s)。仅改变形状的头维度变化可突破 49%;见下方的头维度。按影响程度排序,产生效果的因素如下:
Ironwood XLA 标志 + SparseCore 卸载:将集合操作(all-gather、2D all-gather、reduce-scatter)卸载到 SparseCore,加上一系列针对 v7x 特定的 XLA 标志,使我们的 MFU 从 41% 提升至 44.5%,且损失中性。(完整标志列表见附录 H。)
扩展重材料化:对注意力及 MLP 投影(qkv_proj、q/k/v_proj、out_proj、mlpwi_0、mlpwo、context)进行检查点保存,使其符合激活内存预算;若再增加一个(mlpwi_1),则会导致 HBM 溢出 18 GB。
采用 Splash attention + Tokamax 及 2048-token 块。
在此规模下,分片轴无关紧要:纯 FSDP、4-FSDP×32-DP 和 8-FSDP×16-DP 在 128 个设备上均达到约 1.5 TFLOP/s 的吞吐量;纯 FSDP 因简洁性而胜出。芯片内张量并行(TP=2)总体上是亏损的:在半批次下 MFU 降低 1.6%,在全批次下出现内存溢出(FSDP=64 使每个芯片的权重状态翻倍)。
扩展与缩减,以及为何它是免费的
JAX/XLA 堆栈在此处最有用的特性是
OLMo 3 , developed by the Allen Institute for AI (AI2), is a state-of-the-art, fully open language model trained with a modern architecture and a multi-stage training recipe. To evaluate the capabilities of MaxText on Google Cloud TPUs, our team set out to reproduce AI2’s OLMo 3 7B from scratch. We chose OLMo 3 because it combines three properties that rarely appear together. It is a strong, modern 7B model trained at real production scale. AI2 exposes nearly the complete model flow, including data, code, configurations, checkpoints, logs, and evaluations. And finally, it gives us an independent PyTorch and GPU reference against which we can test MaxText and TPUs.
We reproduced AI2’s OLMo 3 7B in MaxText on Google Cloud TPUs, both the stage-1 pre-training and the stage-2 mid-training anneal, and proved the match on held-out metrics, not just the loss curve:
The main highlights, each covered in detail later in the post:
PyTorch → JAX model conversion. OLMo 3's architecture (reordered-norm block, QK-norm, 3:1 sliding/global attention) ported to MaxText and verified with a logit-parity check: the converted step-0 checkpoint matches the HuggingFace reference at KL ≈ 1.5e-3, the "same model, different framework" noise floor, and at the full 8192-token context in bfloat16 the two agree on the top-1 token 98.75% of the time.
Verification that catches real bugs. Held-out evals caught a data-loader bug that made MaxText look like it was beating the reference; the gain was memorization.
Reliability over a multi-week run. Checkpoint-and-resume replays the run exactly: a controlled A/B shows Δ = 0.000 at every step after a resume, and when a host failure killed the stage-2 run mid-flight, the resumed run re-trained 127 steps at Δ = 0.000 in logged loss and perplexity.
Resizing the training job mid-flight. At step ~1.05M we lost three quarters of our capacity; the run resumed on a slice one quarter the size with no recipe change (same script run_olmo3_7b_stage1.sh scales per device batch size to keep GBS constant), per-device throughput preserved within 1% (≈100% strong scaling, measured in both directions).
Changing TPU generation mid-recipe. Stage 2 pointed the identical launcher at v5p instead of Ironwood, changing only the device type, and sustained 57.4% MFU.
Performance work that paid for itself. 44.5% MFU on Ironwood at 7B via SparseCore collective offload, remat tuning, and optimal sharding, roughly a third of the compute budget bought back.
Co-design for TPU: faster at the same quality. Reshaping attention from 32 heads × head-dim 128 to 16 × 256, at identical parameters and FLOPs, runs +12.4% faster because head-dim 256 fully utilizes Ironwood's 256×256 MXU, and its loss curve matches the original through 120B tokens (30k steps). This was a side ablation; the reproduction kept the original architecture.
Starting from AI2’s step-0 PyTorch weights and the same core recipe, the MaxText run tracks AI2’s published loss curve over the full ~5.93T-token / 1.41M-step budget and lands on top of it at the end of stage-1. We even simplified two recipe details (a single cosine LR schedule where AI2 stitched two, and the publicly released data mix; see t he recipe below), and the match held anyway. The rest of this post is how each of these was built, measured, and, in one instructive case, nearly faked.
Why reproduce OLMo 3?
OLMo 3 is one of the few genuinely open frontier-class language models: open weights, open data, and a fully specified training recipe with a public reference run on Weights & Biases. Matching that independently trained run, on held-out metrics rather than just the loss curve, is strong evidence that the MaxText stack (optimizer, loss, data pipeline, numerics) is faithful, not just “looks like it’s training.”
MaxText is a JAX/XLA LLM training framework built for TPUs. The question we set out to answer: can a PyTorch-on-GPU recipe be reproduced faithfully in JAX-on-TPU, matched on the metrics that matter rather than bit-for-bit, and how do you prove it?
OLMo 3’s recipe is a 3-stage curriculum: general pre-training, mid-training (annealing), and long-context adaptation. This post covers stage 1 (the ~5.9T-token pre-training run) and stage 2 (mid-training), both trained end to end and matched against AI2’s references. Stage 3 and post-training (SFT/RL via Tunix) are recipes we’ve written but not yet run.
OLMo-3 pre-training curriculum: stage 1 (Ironwood) and stage-2 anneal (v5p) reproduced in this post; stage 3 long-context (seq 65k, YaRN) and post-training (SFT then GRPO via Tunix) next
The OLMo-3 curriculum. Stages 1 and 2 are reproduced in this post; stage 3 and post-training are next.
The recipe
OLMo 3 7B is a 32-layer, 4096-dim dense transformer with a few non-standard choices: a “reordered norm” block, QK-norm, and a 3:1 mix of sliding-window and global attention. The MaxText config (olmo3-7b-pt.yml, used for stage 1 and 2) matches it exactly:
The training recipe mirrors OLMo-core’s pretrain-1.py; the knobs that have to match for the curves to line up:
We started training from AI2’s step-0 PyTorch checkpoint , converted to Orbax, so MaxText begins from the exact same weights as the reference. The conversion itself was the first checkpoint: a forward pass on the converted weights matched the HuggingFace reference at KL ≈ 1.5e-3 with 9/10 top-10 token overlap, the “same model, different framework” noise floor.
Does the match depend on inheriting AI2’s initialization? Apparently not. As an independent check we also trained a run from MaxText’s own random init for ~50k steps (3.5% of the horizon); its training loss tracked AI2’s published curve closely, running a touch below it. That’s a training-loss spot check, not a full replicate, but it suggests the match doesn’t hinge on starting from AI2’s weights.
The data pipeline mirrors OLMo-core exactly: tokenize and concatenate all documents (EOS between), slice into non-overlapping 8192-token instances, globally shuffle the index with a fixed seed, and apply an n-gram repetition filter that masks instances with >32 repeated n-grams. MaxText’s dataset_type=olmo_grain (built on Grain ) implements this.
Two deliberate divergences from AI2’s run. (1) LR schedule : AI2 originally planned ~5T tokens and extended the run mid-flight to a final horizon of ~5.93T, so its LR trace stitches two cosine curves (visible in its public WandB run); we ran a single cosine over the full horizon. (2) Data : we train on the publicly released OLMo-3 mix, which omits a small fraction (<0.5% of the token budget, mostly s2pdf shards absent from the released file list) that AI2’s internal run saw. Both are simplifications we chose, not accidents, and MaxText still matches on every held-out surface. This is also why we say “reproduced to within run-to-run noise,” not bit-for-bit (see the KL analysis in §G ).
What we had to build
OLMo 3 wasn't in MaxText when we started; the reproduction added, and upstreamed, everything below. "Reproduce it yourself" at the end of this post is config, not code.
The model itself : the reordered-norm block, QK-norm, and the 3:1 sliding/global attention pattern ( #3004 , #3112 ).
A skip-step optimizer matching OLMo-core's semantics down to the Bessel-corrected running std (skip at 6σ over a 128-step window) ( #3490 ).
z-loss ( #3211 ) and per-parameter weight-decay masking so embeddings can be excluded ( #3280 ).
The olmo_grain data pipeline ( #3749 ): random-access reads of pre-tokenized shards, a seeded global index shuffle with a fingerprint guard against silent data swaps on restart, the n-gram repetition filter, and (from stage 2) Grain iterator-state checkpointing.
HF↔Orbax checkpoint conversion with a logit-parity check, the tool behind every "framework noise floor" number in this post ( #3112 , #3832 ).
The stage-1 and 2 launcher (env-driven run script + XPK wrapper with submit / monitor / resume_until_done) ( #3886 ).
TensorBoard parity tags (optim/step_skipped, perf/total_tokens) so every metric on AI2's W&B dashboard has a MaxText counterpart to compare against.
Does it converge?
The headline is a single overlay: MaxText’s stage-1 lm_loss vs AI2’s published WandB curve, step-aligned and binned to 2k-step means. Through ~800k steps the two track within ±0.012 ; from ~0.9M MaxText edges below AI2 and never crosses back, the first sign of the data bug dissected in the next section.
MaxText vs AI2 stage-1 loss over 1.41M steps, with the 18k-step-smoothed gap in a lower panel
Top: the curves are indistinguishable at this scale until the tail. Bottom: the gap stays inside ±0.012 through ~800k, tilts negative from ~0.9M as data-bug repeats accumulate, and dives past 1.25M, reaching −0.22 in raw 2k-step bins before the 18k-step smoothing both panels are drawn with. None of this reaches held-out loss or downstream accuracy, as the next sections show. (Per-landmark table: Appendix A; full curves committed as olmo_stage1_loss_curve.tsv.)
But a loss curve alone is a weak proof: two runs can match on training loss and diverge on everything you’d actually care about. So we verified convergence on four independent surfaces at six step landmarks spanning 915k steps :
Held-out C4 lm_loss : forward-only eval on 16M tokens of held-out C4-en, identical batches both runs.
8-task lm-eval-harness suite : MMLU, HellaSwag, ARC-easy/challenge, OpenBookQA, PIQA, BoolQ, WinoGrande.
Multi-domain held-out perplexity : a Paloma-style sweep across web, news, encyclopedic, and mixed domains.
Token-level KL : next-token distribution distance on identical inputs.
How we measured. All evals run on step-aligned checkpoint pairs: the live MaxText run vs the AI2 checkpoint at the same step, i.e. its public HuggingFace revision allenai/Olmo-3-1025-7B@stage1-step{N} converted to Orbax (the conversion reproduces the PyTorch reference to KL ≤ 1.8e-3, the framework noise floor). lm-eval uses the standard lm-eval-harness (5-shot MMLU, defaults elsewhere); σ is the per-task harness stderr, combined in quadrature for deltas.
At end of stage-1, every surface agrees the two recipes are interchangeable:
The fourth surface, token-level KL, is the one number that isn’t tiny: mean 0.389 nats on identical inputs, ~200× the framework noise floor. That’s expected for two independent runs of the same recipe (same aggregate skill, different allocation of probability mass), and it’s why we say “run-to-run noise,” not bit-for-bit; the breakdown is in §G .
And the downstream-accuracy gap never exceeds ±0.005 macro across all six landmarks , with the sign flipping four times, exactly the random walk you’d expect from two faithful runs differing only in RNG and numerics (per-landmark table in Appendix D ):
8-task macro accuracy, MaxText vs AI2 across six landmarks, with per-landmark delta
Downstream capability is interchangeable at every landmark. Unlike training loss, accuracy never diverges monotonically; the delta random-walks inside ±0.005 and ends at +0.0002.
The bug that looked like a win
Here’s where it gets interesting. From ~0.9M steps MaxText’s training loss edged below AI2’s and never crossed back; past ~1.25M it pulled clearly under, by −0.06 on average and by as much as −0.25 in a few hundred-step stretches. Watching only the training-loss overlay, you’d conclude MaxText had pulled ahead.
It hadn’t. Held-out C4 loss at the bracketing checkpoints was tied (Δ −0.004 at 1,000k, +0.003 at end of stage-1), and downstream accuracy at 1,350k slightly favored AI2 . Training loss was dropping while generalization didn’t move. That’s the signature of memorization : the model was seeing some sequences more than once and scoring low loss on the repeats.
Top: MaxText training loss dives to 1.63 while AI2 stays flat. Bottom: held-out C4 delta stays near zero throughout
The whole story in one figure. Top: in the 1.24M–1.41M window MaxText’s training loss repeatedly plunges to 1.63 on 2k-step means (Δ −0.22; −0.25 in hundred-step stretches) while AI2’s stays flat. It looks like a runaway win. Bottom: the training-loss Δ (blue) drifts negative, but held-out C4 loss (green diamonds) never leaves the ±0.02 band. Training loss dropped; generalization didn’t.
The cause was a double-sharding bug in the Grain data loader . MaxText’s OLMo loader passed ShardOptions(shard_index, shard_count) to the Grain DataLoader while the index sampler was already sharding internally. Grain’s shard_options doesn’t just record metadata; it re-strides the sampler’s index stream. With shard_count=32, the data cursor advanced 32× too fast , so stage-1 stopped being one clean epoch and became a Poisson(≈1) resample-with-replacement: roughly 37% of the corpus never seen, 37% seen once, 26% seen twice or more. The token budget was unchanged (~5.9T real tokens), which is why the loss still tracked AI2 globally, but the repeated instances deflated training loss exactly where they recurred.
Two lessons came out of this:
Training loss is not a convergence proof. The only reason we didn’t ship a false “MaxText beats the reference” claim is that we’d committed to held-out eval at every landmark. The memorization dip is invisible on held-out C4 and on all 8 downstream tasks.
Honest reproductions need a bug budget. The fix (grain.sharding.NoSharding(), letting the sampler own all sharding) is one line. Finding it took an A/B harness, a unit test that reproduces the divergence at shard_count>1, and a hardware re-run to validate.
While validating the fix we found a second , independent bug: an off-by-one in resume-step detection. The checkpoint directory number is N, but the train loop writes dir N after iteration N completes, so the model restored to step N+1 while the data loader resumed at batch N, re-training one batch and then running permanently one step behind. With both fixes, a checkpoint-and-resume run replays the uninterrupted run exactly: Δ = 0.000 in logged loss at all 99 steps . (That A/B ran single-worker data loading; the multi-worker case surfaced in stage 2, where we closed it; see Stage 2 .) Both bugs have regression tests that fail on the old code and pass on the fix.
We let the in-flight stage-1 run finish as-is: it was 85% done, the fix can’t un-scramble already-read data, and a relaunch would forfeit ~1.2M steps of compute. The verification above shows the bug cost zero observable accuracy ; the fix is for future runs.
Performance and scale on Ironwood
Reproducing the math is half the job; the other half is making it fast, and keeping it fast when the cluster shifts under you. Over the weeks the 1.4M-step run took, the job was preempted, rescheduled, and resized more than once, and the stack had to absorb all of it without touching the recipe.
Squeezing out MFU
On Ironwood at 7B, per-device batch 4, we landed at 44.5% MFU (510–513 TFLOP/s/device) for the stock architecture (“variant D,” our label from the ablation sweep in Appendix J ). A shape-only head-dim change clears 49%; see Head-dim below. What moved the needle, in order of impact:
Ironwood XLA flags + SparseCore offload : offloading collectives (all-gather, 2D all-gather, reduce-scatter) to the SparseCore, plus a set of v7x-specific XLA flags, took us from 41% to 44.5% MFU, loss-neutral. (Full flag list in Appendix H .)
Extended rematerialization : checkpointing the attention and MLP projections (qkv_proj, q/k/v_proj, out_proj, mlpwi_0, mlpwo, context) fit the activation-memory budget; adding one more (mlpwi_1) overflowed HBM by 18 GB.
Splash attention + Tokamax with 2048-token blocks.
Sharding axis is irrelevant at this scale : pure FSDP, 4-FSDP×32-DP, and 8-FSDP×16-DP were all within ~1.5 TFLOP/s on 128 devices; pure FSDP wins on simplicity. Intra-chip tensor parallelism (TP=2) was a net loss: −1.6% MFU at half batch, OOM at full batch (FSDP=64 doubles per-chip weight state).
Scaling up and down, and why it was free
The single most useful property of the JAX/XLA stack here is that the
| 刊期 | 得分 | 排名 | 结果 |
|---|---|---|---|
| 2026-09-28 | 8.4 | 19 | 入选 |
| 2026-09-27 | 8.4 | 30 | 未入选 |
| 2026-09-26 | 8.4 | 43 | 未入选 |
| 2026-09-25 | 8.4 | 52 | 未入选 |