機器學習系統與基礎設施
分散式檢查點(distributed checkpointing)
一次大型訓練可能把數兆位元組的參數與最佳化器狀態散在數千張 GPU 上,而它必須週期性地把這一切存下來,才不會讓一次故障抹掉好幾天的進度。若每個分片都擠過單一行程、寫單一檔案,存檢查點就會卡住整個叢集。分散式檢查點改成讓每個 rank 平行地把自己那一片狀態直接寫進儲存。
每個 rank 序列化它擁有的分片並同時寫出,產生一個分片檢查點,外加記錄全域張量如何被切分的中介資料。因為那份中介資料是邏輯性的、並未綁死在固定的行程數上,檢查點在載入時可被重新分片(re-shard)——還原到不同數量的 GPU 或不同的平行佈局上,這正是 PyTorch DCP 之類系統提供的能力。效能工作的重心在於讓存檔與運算重疊:先快速把狀態複製到主機記憶體或暫存緩衝區,再趁訓練繼續時非同步刷寫到平行檔案系統或物件儲存。總寫入頻寬必須隨模型與叢集規模擴展,否則存檢查點的時間就會反客為主。
頻繁而便宜的檢查點是容錯訓練的根基——檢查點越省,你就能存得越勤,一次故障所損失的工作也越少。可重新分片的檢查點還把「你怎麼訓練」與「你怎麼續訓」解耦,這在彈性排程器於訓練途中更動世界大小(world size)時格外重要。
又稱
另見