你無法在完整規模下調超參數
前沿訓練是一次性的:你只有一個學習率、一套初始化方案,若它們錯了,你可能要等到燒掉一大筆算力後才發覺。靠掃描做的一般 超參數調校 行不通,因為每次試驗的成本都與最終模型相當。這正是 最大更新參數化(muP) 所解決的問題:它以有原則的方式,讓初始化變異數與各層學習率隨模型寬度縮放,使最佳超參數在模型放大時保持不變。
- 在 muP 下重新參數化模型,使激活值、梯度與更新量在寬度增長時都維持 order-1 的尺度。
- 在一個小型但已 muP 參數化、訓練便宜的代理模型上,掃描學習率、暖身與初始化。
- 把勝出的超參數直接移轉到完整寬度的模型——不必在大規模下重新調校。
在 muP 下,最优学习率与宽度无关——在小模型上调好后直接迁移到全尺度训练。
延伸模型而不從頭再來
一旦基礎模型存在,你很少為了加入一項能力而從頭重訓。持續預訓練 從既有檢查點接續訓練於新資料——一個全新領域(程式碼、醫學、一種新語言)、更長的上下文視窗,或單純更多同類資料以把損失曲線再往下推。其中的技藝是避免災難性遺忘:在新語料上天真地用高學習率,會抹去模型原本已知的東西。常見配方是溫和地重新暖身的學習率、一個保留部分原始分布切片的資料組成,以及一個大到足以學會新領域、但相對於原始預訓練仍偏小的詞元預算。
持續預訓練也是前面各篇的 FP8 與平行化選擇被整批重用之處:你接續進入相同的三維佈局,常常 FP8 訓練 仍然開著,唯一的新麻煩是確保優化器狀態與學習率排程能連貫地接續,而非從頭冷啟。
在一萬張 GPU 上,總有東西壞掉
硬體可靠度是殘酷的算術。若單張 GPU 的平均故障間隔是數年,一萬張的叢集每隔幾小時就會遇到一次故障。一場無法在某個節點死亡後存活的訓練,就是一場永遠跑不完的訓練。容錯訓練 是一門預期故障的紀律:偵測死掉的 rank、把它逐出、再從最後一個良好檢查點重啟,並把損失的工作量降到最低。整套策略都建立在檢查點要夠快之上——因為檢查點越貴,你能負擔的存檔頻率就越低,回滾時損失的進度也越多。
集群可靠性是残酷的算术:一万块 GPU 把单卡数年的寿命压缩成每隔几小时就有一次故障。
跟得上叢集的檢查點機制
當狀態分片在數千個 rank 上時,你無法把它全擠過單一寫入者。分散式檢查點 讓每個 rank 平行寫出自己的那一片——這就是 檢查點分片——使存檔的牆鐘時間在叢集擴大時大致保持不變。微妙的要求是,存下的格式必須懂得重新分片(resharding-aware):你應該能把在某種平行佈局(例如 TP=8、PP=16)下寫出的檢查點,在以不同形狀的叢集恢復時,重新載入到另一種佈局——這意味著檢查點必須儲存邏輯上、與佈局無關的張量,而非每個 rank 的原始實體分片。
- 非同步存檔:快速把張量快照到主機記憶體,再於背景在訓練繼續的同時刷寫到儲存裝置。
- 儲存與佈局無關的邏輯張量,使檢查點能重新分片到不同的平行組態。
- 依預期故障率設定檢查點間隔:頻繁到足以限制損失的工作量,又稀疏到不至於主宰整場訓練。
- 在重新載入時驗證——一個直到重啟才發現損毀的檢查點,本身就是另一種災難。
Young–Daly 最优公式根据检查点开销 δ 和集群平均故障间隔 M 来确定检查点间隔。