產業脈動 2026 年 9 月 25 日

2026-09-25 — MaxText 在 TPU 上重現 OLMo 3 7B 預訓練全過程

primary=https://developers.googleblog.com/en/reproducing-olmo-3-7b-pre-training-in-maxtext-case-study-of-large-scale-training-on-tpus/ primary=https://github.com/AI-Hypercomputer/maxtext

MaxText 在 TPU 上重現 OLMo 3 7B 預訓練全過程

Google Developers Blog · 2026-09-24

Google 的 MaxText 團隊把 AI2 用 PyTorch 在 GPU 上訓練的 OLMo 3 7B,整套配方搬到 JAX 上、換成 TPU 重跑一次,兩條訓練損失曲線幾乎疊在一起。這篇案例文發表於 2026 年 9 月 24 日,主要在 Ironwood(256 顆晶片、4×8×8 切片)與 TPU v5p-256 上完成兩階段訓練。

原本的問題

跨框架重現一個訓練配方,過去只能看訓練損失曲線像不像——曲線貼合就算「重現成功」。但訓練損失本身會被資料管線的 bug 蒙騙:只要某些樣本被重複餵進模型,loss 會看起來更低,卻不代表模型真的學得更好、更接近原始模型的行為。業界很少有機會拿到一個模型完整的資料、程式碼、config、checkpoint 與評測紀錄,又同時有獨立的 PyTorch/GPU 版本可以對照,所以這類「換框架、換硬體後是否還是同一個模型」的問題,長期缺乏可驗證的答案。

採用的方法

團隊沒有只看訓練損失,而是在 915k 步範圍內的六個步數里程碑,同時檢查四種獨立指標:16M token 的 held-out C4 語言模型損失、MMLU/HellaSwag/ARC 等八項任務的 lm-eval-harness 準確率、Paloma 風格的多領域 held-out perplexity,以及相同輸入下的 token 級 KL 散度。訓練面則採用純 FSDP 切分(4-FSDP×32-DP 與 8-FSDP×16-DP 效能差在 1.5% 內)、對 qkv_proj、out_proj、mlpwi_0 等層做重算式 checkpoint,並確認 optimizer moment 必須用 float32——換成 bfloat16 會在 1000 步內多墊出 0.93 的 loss 落差。

項目Stage 1Stage 2
硬體Ironwood 256 顆(4×8×8)TPU v5p-256(128 顆)
吞吐510–513 TFLOP/s/裝置,MFU 44.5%263 TFLOP/s/顆,MFU 57.4%
訓練量~1.41M 步、~5.93T tokens47,684 步、~100B tokens
等效耗時~77k chip-hours(約 12.5 天)~5k chip-hours(約 39 小時)
batch全域 512、每步 4.19M tokens,序列長度 8,192

實際效果

Stage 1 的 held-out C4 損失與原始 AI2 版本相差僅 −0.004 nats,八任務下游準確率差 +0.0002,相同輸入下的 token 級 KL 散度均值 0.389 nats(約為框架雜訊底線的 200 倍量級,仍在可接受範圍)。Stage 2 的 held-out C4 損失穩定維持在約 +0.006 nats、下游準確率差 −0.0023,都落在評測雜訊範圍內。擴縮測試把裝置數從 128 拉到 512(4 倍),聚合吞吐拿到 3.99 倍,幾乎線性;在第 ~1.05M 步把切片縮回 128 顆,單裝置吞吐仍維持 510–513 TFLOP/s,不用改配方。

過程中也真的抓到一個訓練損失掩蓋不了的問題:Grain 資料載入器的雙重分片 bug 讓約 37% 的語料從未被看過、26% 被重複取樣兩次以上,訓練損失因此被壓低,但 held-out 損失與下游準確率完全沒動——這正是四指標交叉驗證抓出來、單看訓練損失會漏掉的案例。訓練損失曲線在前 ~800k 步都貼在原始曲線 ±0.012 以內,直到約第 0.9M 步才開始出現這個由資料 bug 造成的偏移。checkpoint 續訓的精確度也一併驗證:Stage 1 在修復後 99 個步數的續訓誤差都是 0,Stage 2 因主機故障重跑的 127 步,loss 與 perplexity 同樣對到 0 誤差。

團隊也做了一項架構層面的消融:把原本 32 個 head、每個 128 維的注意力設定,改成 16 個 head、每個 256 維,在不改變模型行為的前提下把單裝置吞吐從 508 拉到 571 TFLOP/s,提升 12.4%。這說明重現一個模型不代表要照抄每一個超參數,只要驗證方法夠嚴謹,仍有調整硬體友善度的空間。整條 pipeline 另外把 all-gather、2D all-gather、reduce-scatter 等集合通訊卸載到 SparseCore,讓 FSDP 通訊不占用主計算路徑。

對正在用 MaxText 或其他框架把既有 PyTorch/GPU 訓練配方搬去 TPU/JAX 的團隊,這篇案例文點出的具體檢查項是:不要只憑訓練損失曲線判斷「重現成功」,要在固定的步數里程碑上同時看 held-out loss、下游任務準確率與 KL 散度;多主機 resharding 用的資料載入器要特別檢查是否有重複取樣的分片 bug;optimizer moment 精度用 bfloat16 前,先確認對 loss 的影響量級。

原始來源:Google Developers Blog、MaxText GitHub


End of article
0
Would love your thoughts, please comment.x
()
x