← 返回 2026-09-28 简报

在MaxText中复现OLMo 3 7B预训练:TPU大规模训练案例研究

Reproducing OLMo 3 7B Pre-training in MaxText: case study of large scale training on TPUs

语音播报
摘要
事件:Google团队在TPU上复现AI2的OLMo 3 7B模型,验证MaxText框架的忠实度与可靠性。 要点:通过KL散度和保留指标匹配参考曲线;实现断点续训零偏差及跨代际TPU无缝切换。 影响:证明JAX/TPU栈可高保真复现PyTorch/GPU训练,为大规模开源模型高效部署提供实证。

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 堆栈在此处最有用的特性是

本条评分 8.4 score-v1
  • 来源权威 8
    注册表 priority=8(Google Developers Blog)
  • 时效 0.4
    没有发布日期(nodate 源),给地板分 0.4(证明不了它新,不按新的加分)
  • 多源印证 0
    只有 1 家在报(无旁证)
  • 社区信号 0
    无社区数据(本管线走 RSS,HN 的 hn_fetcher 未接入)
历史
刊期得分排名结果
2026-09-28 8.4 19 入选
2026-09-27 8.4 30 未入选
2026-09-26 8.4 43 未入选
2026-09-25 8.4 52 未入选
原文链接:https://developers.googleblog.com/reproducing-olmo-3-7b-pre-training-in-maxtext-case-study-of-large-scale-training-on-tpus/
来源:Google Developers Blog
以上内容由 AI 自动翻译,仅供参考。
← 返回简报