檢查點分片和可恢復訓練
將模型的訓練狀態保存為片段(分片)的技術,以便可以保存和重新加載大型模型,而不會受到內存或磁碟限制的影響,因此崩潰的運行可以準確地從中斷的地方繼續。
概述
對於任何需要在多張 GPU 上持續數天甚至數週的訓練工作來說,這也是必不可少的。
深入探討
訓練檢查點是恢復所需的所有內容的快照:模型權重、優化器狀態、學習率計劃、資料載入器的位置和隨機數產生器種子。對於大型模型,此快照可能有數百 GB,對於單一檔案或單一電腦的記憶體來說太大了。檢查點分片將該快照拆分為多個檔案和多個等級,因此每個 GPU 僅並行寫入自己的切片。然後,可恢復訓練會重新載入這些分片並精確地恢復完整狀態。如果沒有它,在第 200 小時崩潰的多周運行將不得不從頭開始。 PyTorch 分散式檢查點、DeepSpeed 和 Hugging Face Hub 的分片安全張量格式等框架可以實現此例程。
技術洞察
分片之所以有效,是因為分散式訓練已經跨等級劃分了權重和優化器狀態(透過資料、張量或零並行)。每個等級僅序列化其分區,通常序列化為安全張量之類的格式,允許延遲、記憶體映射載入。索引檔案將參數名稱對應到分片檔案。為了確定性地恢復,系統還保留 RNG 狀態、優化器步數和確切的資料載入器偏移量,因此重新運行會重現相同的批次序列。
戰略影響
成本與預算
多年來,架構決策決定著效能和營運成本。
更明確的決策
技術教育幫助團隊選擇正確的堆疊,而不僅僅是最新的堆疊。
品質管控
更好的工程選擇可以減少生產中的可靠性事故。
檢查點分片和可恢復訓練的未來
檢查點正在從週期性的停止世界事件轉變為非同步且幾乎免費的事件。預計會有更多的記憶體中和重疊檢查點,在訓練繼續的同時在後台寫入分片,加上糾刪碼和複製檢查點,可以在千個 GPU 規模上常見的節點故障中倖存下來。雲端物件儲存和更快的本地 NVMe 層將託管分片,而安全張量等標準化格式將不斷改進訓練復原和推理部署的安全性、快速、部分載入。
現實世界的實施
前沿模型運行在數千個 GPU 上,每隔幾百步自動保存分片檢查點,因此單一失敗的節點只需要幾分鐘,而不是幾天。
Hugging Face 將大型開放式模型分發為多個 safetensors 分片和一個 index.json,以便用戶可以逐一下載和載入它。
研究人員恢復中斷的微調,恢復精確的優化器動力、步數和資料載入器位置以無縫繼續。
在廉價的可搶佔式雲端 GPU 上進行現貨實例訓練,其中頻繁的分片檢查點使作業能夠在被驅逐和重新安排的情況下倖存下來。
風險與防護欄
優化一項基準測試可以隱藏更廣泛的系統弱點。
基礎設施和維護成本常常被低估。
隨著系統變得更加複雜,安全性和可觀察性差距可能會擴大。
實施路線圖
在實施之前定義延遲、品質和成本目標。
在實際負載和資料條件下進行基準測試。
儀器監控錯誤、漂移和使用者影響。
在擴展之前準備回滾和事件回應路徑。
不斷探索
Free newsletter
Get the daily AI briefing
Three verified AI stories every weekday morning, written in plain English. Free forever, no ads.
One email each weekday. Unsubscribe in one click. We never sell or share your address.
Test yourself
Take the Checkpoint Sharding and Resumable Training quiz
Instant feedback on every answer, and a shareable certificate with a verifiable ID once you pass a course.
Support free AI education. AI Understanding is a 501(c)(3) nonprofit — no ads, no paywall, ever. Make a donation
常見問題
什麼是檢查點分片與可恢復訓練?
將模型的訓練狀態保存為片段(分片)的技術,以便可以保存和重新加載大型模型,而不會受到內存或磁碟限制的影響,因此崩潰的運行可以準確地從中斷的地方繼續。對於任何在多個 GPU 上運行數天或數週的訓練作業來說都是至關重要的。
為什麼大模型檢查點要分成分片?
將巨大的檢查點(數百 GB)作為一個檔案是不切實際的;分片允許許多佇列並行寫入其切片並啟用分段載入。
除了模型權重之外,可恢復檢查點通常還必須保存什麼?
為了準確地恢復,檢查點儲存優化器狀態、隨機數產生器狀態、步數以及資料載入器停止的位置。
長期工作可恢復培訓的主要好處是什麼?
透過可恢復的檢查點,中斷的運行會重新載入其上次保存的狀態並繼續,而不是丟棄數天的計算。
每個 GPU 等級通常如何對分片檢查點做出貢獻?
由於分散式訓練已經對模型進行了跨等級劃分,因此每個等級僅寫入其切片,從而實現並行保存和記憶體高效。
索引檔案在分片檢查點中扮演什麼角色?
索引(例如,index.json)記錄哪個分片保存每個參數,以便載入器可以取得並重新組裝正確的片段。