你的 GPT-6 訓練任務已經在 AI 訓練叢集上跑了 10 天,卻突然當機了。螢幕熄滅,心跳加速。之前的進度難道全沒了嗎?想要恢復訓練進度,必須事先規劃。沒有人願意眼睜睜看著幾週的算力投入付諸流水,也沒有人願意面對漫長停機帶來的壓力。

這份實戰指南將提供切實可行的步驟。首先,弄清楚哪些狀態遺失了。然後,儲存並恢復這些狀態以繼續訓練。接著,調整學習率。最後,驗證恢復是否成功。

理解被中斷的訓練工作階段

AI 訓練叢集故障的常見原因

訓練任務中斷的原因有很多,其中硬體故障最為常見。GPU 可能因過熱而自動關機,記憶體模組可能出現錯誤,電源也可能在毫無預警的情況下失效。軟體缺陷同樣會導致當機。一個細微的程式碼錯誤,可能在執行數天後才破壞訓練迴圈。網路問題也可能讓節點與叢集失去連線。無論是哪種故障,結果都一樣:你的訓練任務會在中途停止。

其影響遠不只任務停止本身。你的叢集會處於閒置狀態,其他任務在佇列中等待,團隊的有效工作時間被浪費。這樣的停機損失不只是算力成本,更會打亂研究節奏,拖慢專案進度。理解這些原因,有助於你為不可避免的中斷做好準備。你無法阻止所有故障,但你可以掌控自己的應對方式。

訓練停止時會遺失哪些狀態?

很多工程師誤以為只需要保存模型權重,這種想法往往會導致恢復失敗。真正的續訓,要求你恢復完整的訓練狀態。模型參數只是其中一部分,最佳化器狀態同樣至關重要。以 Adam 最佳化器為例,它保存了每個參數對應的動量與變異數估計值。沒有這些值,最佳化器的行為就會像剛開始訓練一樣。學習率排程器也保存了位置信息,它知道目前訓練進行到衰減計畫的哪一步。

隨機數產生器(RNG)狀態同樣重要。訓練過程中,資料打亂與資料增強都依賴隨機性。每個 epoch 的資料順序都會不同,影像變換會隨機裁切與翻轉。這些操作都由 RNG 狀態決定。如果不恢復它,資料管線就會產生不同的樣本序列。這種不一致會帶來細微但真實的訓練差異,模型的收斂路徑可能因此改變,驗證結果也可能與當機前的日誌不再一致。

一個完整的檢查點應當同時保存這些組成部分:模型、最佳化器、排程器以及 RNG 狀態。採用這種完整保存方式,恢復過程就會簡單得多。你可以把所有內容恢復到故障前的精確狀態,讓訓練彷彿從未中斷。這樣的準備,能夠把原本可能是一場災難的故障,變成一次小小的不便。你為完善檢查點機制投入的時間,會在今後的每一次故障中獲得回報。

透過檢查點恢復訓練進度

保存完整的模型與最佳化器狀態

檢查點必須捕捉訓練狀態中的每一個關鍵部分。你不能只保存模型權重。最佳化器保存了每個參數的動量值,排程器記錄了自己在學習率衰減曲線中的位置,隨機數產生器決定了資料打亂的模式。每一個元件,都是平順恢復訓練不可或缺的一環。

你可以透過一條命令建立一個完整的檢查點。字典結構能讓所有內容保持清楚有序。保存目前 epoch 編號,以便知道從哪裡重新開始;保存模型的狀態字典,以還原網路架構中的參數;保存最佳化器的狀態字典,以保留已經累積的梯度資訊;保存排程器狀態,以延續既定的學習率計畫;記錄 RNG 狀態,以保證資料管線的可重現性。這個完整快照,能在任何故障後真正實現訓練進度的恢復。

torch.save({
    'epoch': epoch,
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
    'scheduler_state_dict': scheduler.state_dict(),
    'rng_state': torch.get_rng_state()
}, 'checkpoint.pth')

隨著模型規模增大,這個檔案會迅速膨脹。一個簡單經驗法則是:若包含最佳化器狀態,每個參數大約需要 12 位元組。權重本身占 4 位元組,最佳化器狀態占 8 位元組。這個估算規律可以幫助你在訓練開始前就規劃好儲存容量。

儲存檢查點以實現快速恢復

你選擇的儲存媒介,將直接決定恢復所需時間。像 HDFS 或 S3 這樣的分散式檔案系統提供更強的持久性,而本機 NVMe 則提供更高的速度。兩者各有用途,也往往都需要。活躍訓練過程通常寫入平行檔案系統,長期封存則轉移到物件儲存中,其每 TB 成本通常只有前者的十分之一。

檢查點保存間隔需要仔細權衡。保存過於頻繁,可以減少當機後需要重算的工作量;但保存過於頻繁,也會增加寫入時對 GPU 訓練的阻塞。你必須在兩者之間找到平衡。每 5 分鐘保存一次通常過於激進,而對於持續數天、故障並不頻繁的訓練任務來說,每小時保存一次往往更合適。

檢查點保存間隔應與任務的當機率(平均中斷時間)成比例,而這個中斷時間通常會隨著所使用 GPU 數量的增加而縮短。

大規模訓練任務在實務中通常遵循這一原則。一個 4050 億參數的模型,如果平均 150 分鐘會被中斷一次,就會每 15 分鐘保存一次檢查點;一個 8000 億參數的模型,則可能每 40 分鐘保存一次。通常可將保存間隔設為預計故障間隔的大約十分之一,這樣既能減少算力浪費,又不會帶來過高額外開銷。

現代大型語言模型每次寫入檢查點時,往往需要寫出 350–500 GB 的模型狀態。為了盡量減少訓練中斷,寫入過程必須在 5 分鐘內完成。這意味著系統需要提供 1–2 GB/s 的持續突發寫入頻寬。因此,儲存架構必須優先考量突發效能,而不僅僅是長期吞吐能力。

像 TrainMover 這樣的系統,進一步提升了恢復能力。替換機器會在真正發生故障前,預先完成 CUDA 核心編譯並建立 NCCL 通訊群組。這種預熱機制把檢查點載入排除在關鍵恢復路徑之外。其基於增量更新的設計,只更新離開節點與加入節點之間的連線,其他所有機器都無須修改。即使在 1,024 張 GPU 的規模下,這種架構也能將停機時間基本穩定在 20 秒以內。與此同時,由於準備狀態保存在 CPU 記憶體或 NVMe 中,因此不會占用任何 GPU 顯示記憶體,GPU 顯示記憶體在最終切換前始終保持不變。

非同步檢查點是另一種有效的緩解策略。你可以先將資料從 VRAM 卸載到主機記憶體中,然後讓 GPU 恢復訓練,而由 CPU 在背景負責保存。對 85,000 次檢查點的分析顯示,這種方式可以將全域頻寬控制在 1 TB/s 以下,而本機卸載速度可達到 50–200 GB/s。這樣就能避免保存過程中 GPU 長時間停頓。

對於大規模訓練而言,檢查點期間 GPU 閒置帶來的成本每天可能超過 4,000 美元。最佳化停機時間,實際上就是在直接降低成本。你每減少一秒恢復時間,整個叢集都會隨之節省開支。因此,檢查點策略應當與模型架構同等重視。你在穩健保存機制上投入的精力,會在之後的每一次故障中持續帶來回報。

恢復 AI 訓練叢集狀態

載入檢查點並繼續訓練迴圈

恢復工作的起點,是一次簡單的載入操作。你把之前保存的檢查點檔案重新讀入記憶體,這一動作會恢復你先前保存的每一個元件。模型權重回到故障前的數值,最佳化器重新獲得它的動量估計,排程器恢復到學習率衰減曲線中的原有位置。整個訓練環境會回到當機前的精確狀態。

checkpoint = torch.load('checkpoint.pth')
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
start_epoch = checkpoint['epoch'] + 1

其中最關鍵的細節在最後一行:你必須從已保存 epoch 的下一個 epoch 開始,而不是從零重新開始。訓練迴圈也需要相應調整,以便遵循這個起點。只需修改迴圈範圍,使其從恢復後的 epoch 開始,就能避免重複計算已完成的輪次。你的任務將從中斷處精確續上,保留此前所有進度。

for epoch in range(start_epoch, total_epochs):
    train_one_epoch(model, train_loader, optimizer, scheduler)
    validate(model, val_loader)
    save_checkpoint(model, optimizer, scheduler, epoch)

像 fastai 這樣的框架,還能把這個過程進一步簡化。fit_one_cycle 函式會套用循環式學習率與動量策略。如果中斷後處理不當,整個循環會從頭開始,導致結果與未中斷訓練不同。fastai 透過 start_epoch 參數解決了這個問題。你只需重新實例化 learner,並用下一個 epoch 編號呼叫 fit_one_cycle。框架會自動載入此前為該 epoch 保存的檔案,無須手動重新載入權重。這樣,學習率和動量循環就能從準確位置繼續執行,訓練策略得以無縫延續。

重新建立資料載入器與隨機種子

如果資料管線不能一致,僅僅恢復模型也沒有意義。隨機數產生器狀態決定了資料載入器如何打亂樣本順序。每個 epoch 都會產生不同的順序,而影像增強中的各種隨機變換也都依賴 RNG 狀態。

恢復這個狀態其實很簡單。只需把之前保存的 RNG 值重新傳給 PyTorch,資料載入器就會產生與未當機時完全相同的樣本序列。這種可重現性對於驗證一致性非常重要。你的損失曲線能夠在中斷前後保持可比,梯度模式也會更加穩定可預測。

torch.set_rng_state(checkpoint['rng_state'])

你還必須恢復各個 worker 專屬的隨機種子。資料載入器通常會啟用多個 worker 行程,而每個 worker 都維護著自己的隨機狀態。你需要在重新建立資料載入器實例之前,先設定這些種子。這樣才能確保每個 worker 都重現出原本的資料打亂模式。至此,你的訓練環境才算真正回到了故障前的狀態。

最後的驗證步驟,用於確認恢復是否真的成功。先跑幾個 batch,並把損失值與當機前的日誌進行比對。兩者應當大致接近,梯度範數也應落在合理範圍內。如果偏差明顯,就表示恢復過程存在問題。及早發現這些錯誤,能避免你在損壞狀態下繼續浪費數小時計算資源。

你的檢查點載入策略,也會直接影響恢復時間。一個結構清晰、組織良好的檢查點,能讓你迅速定位到正確檔案,並順利完成載入,無須複雜解析。這樣,你的任務通常可以在幾分鐘內恢復到完整執行狀態。高效率的恢復機制,會把一場本可能很嚴重的故障,變成一次短暫的暫停,讓訓練以最少停機、最大把握繼續推進。

調整學習率以繼續訓練

學習率排程器保存著訓練進展中的關鍵資訊。它記錄了目前在衰減序列中的位置。僅恢復模型和最佳化器,並不表示恢復工作已經完成,排程器狀態同樣必須回到當機前的位置。如果你新建一個排程器,整個學習率計畫就會被打亂,因為模型此時本應使用後期更低的學習率。

計算正確的 step 和 epoch

排程器通常根據 step 或 epoch 來決定目前應使用的學習率。你必須知道當機發生時所處的精確位置。可以使用以下公式:current_step = epoch * len(train_loader) + batch_index。這個計算結果給出了你在學習率計畫中的精確 step 編號。然後,將 last_epoch 參數設定為這個 step 值,排程器就能從那個準確位置繼續運作。學習率曲線將跨越中斷點保持平滑,模型在後續每一步都能獲得正確的學習率。這種做法可確保恢復過程不會破壞原本的訓練節奏。

對於多週期排程策略來說,epoch 編號同樣重要。餘弦退火(cosine annealing)和 one-cycle 策略都依賴於完整週期長度。你的檢查點中保存了最後一個已完成的 epoch,只需在此基礎上加一,再傳入排程器建構函式,排程器就能正確計算剩餘衰減過程。這一步能讓恢復過程避免憑感覺猜測。

重新初始化排程器以保持連續性

訓練當機後,不要新建一個全新的排程器。新的排程器會把學習率重設為初始值,而此時模型通常已經進入訓練後期,理應使用更低學習率。突然回到較高學習率,可能導致梯度不穩定、損失暴漲,甚至權重發散。這樣的錯誤,會讓你之前的恢復努力前功盡棄,甚至迫使你再次重新啟動訓練。

如果在當機後重新建立一個新的排程器,你的學習率計畫就會從頭開始。這個突增可能讓模型不穩定,並抵銷數小時的訓練成果。正確做法始終是恢復已保存的排程器狀態。

正確方式是先建立排程器,再載入其保存的狀態。scheduler.load_state_dict() 方法會恢復排程器內部所有參數,讓它回到原本的 step 計數位置。這樣,學習率就能沿著既定軌跡繼續前進,彷彿故障從未發生。整個訓練過程將無縫銜接,不錯過任何一步。

妥善恢復排程器,可以徹底規避上述風險。你的訓練動態得以延續,驗證指標也保持可比。應當把這種做法套用到每一個檢查點中。保存排程器狀態幾乎不需要額外成本,而忽略它的代價卻可能非常高昂。正是這種對細節的重視,才能把一次當機從重大挫折,變成短暫暫停。

驗證恢復是否成功

對損失與梯度進行合理性檢查

在恢復檢查點之後,先執行少量 batch。把此時的損失值與當機前日誌中的數值進行比較。損失應與該 epoch 最後記錄的值大致接近。如果匹配良好,就表示訓練狀態恢復正確;如果出現明顯跳變,則意味著檢查點載入流程存在問題。

你還需要檢查梯度範數。這些數值衡量的是權重更新的幅度。把它們與當機前的日誌進行對照,若結果相近,通常表示恢復成功;若差異明顯,則表示某處出了問題。例如,排程器狀態可能沒有正確恢復,或者隨機數產生器與原始環境不一致。這些細微的不匹配,都會導致訓練結果不穩定。

如果損失值或梯度數值出現大幅跳變,就表示恢復過程出現了錯誤。檢查點檔案可能已損壞,也可能內容不完整。這種快速檢查能幫助你及早發現問題,避免在錯誤狀態下繼續訓練數小時,從而節省寶貴算力。

建議採用對照式驗證方法。提前記錄當機前某一步的損失值與梯度範數,恢復後使用相同數量的 batch 重新跑一遍,再記錄新的結果,並進行逐項對比。兩組數值應當高度接近。整個過程只需要幾分鐘,卻能在後面節省大量時間。

監控恢復初期是否出現發散

恢復後的最初幾個 step,往往最能說明問題。理想情況下,訓練應當和中斷前完全一致地繼續進行,損失曲線保持原有下降趨勢,梯度模式也應維持穩定。

你需要特別留意一些警告訊號。比如梯度爆炸,會表現為參數更新突然異常增大;NaN 值則意味著出現了數值不穩定。這兩種情況通常都源於恢復錯誤。如果不及時處理,訓練不會自行回到正常狀態,反而會繼續浪費資源。因此,必須及早發現。

如果觀察到發散現象,應立即停止任務,不要抱著僥倖心理繼續執行。請檢查檢查點檔案是否完整,確認各個元件是否都已正確載入,並核對 RNG 狀態是否與原始環境一致。

一次乾淨、正確的恢復,不會出現發散。損失曲線會在各個 epoch 之間平滑延續,梯度範數保持穩定可預測,學習過程繼續正常推進。這表示你已成功從故障中恢復,任務能夠按計畫繼續執行,停機階段正式結束,訓練進度也得以延續。

至此,整個驗證流程完成。此後,你就可以對模型後續結果保持足夠信心。

現在,你已經理解了完整的恢復流程:保存全部狀態——模型、最佳化器、排程器和 RNG;實施穩健且間隔合理的檢查點機制;精確恢復訓練環境;正確延續學習率計畫。這些步驟,能夠把伺服器故障從災難性事件,變成可控的短暫中斷。

只要準備充分,你的訓練任務就能在當機後繼續存活。前期在檢查點機制上的投入,會在未來節省無數小時。恢復時間將從數天縮短到幾分鐘。你的訓練會從中斷處準確接續,不浪費算力,也不遺失進度。

你保存下的每一個 epoch,都是對停機風險的保險;你做出的每一次學習率修正,都是對訓練穩定性的保障;你寫下的每一個檢查點,都是對模型軌跡的守護。

有了這份指南,你就能更從容地面對下一次伺服器故障,因為你知道訓練進度是安全的,接下來的道路也是清晰的。請在下一次當機發生之前,就把這些實務做法落實到位,而不是事後補救。

常見問題

訓練過程中應多久保存一次檢查點?

你應根據預期當機率來設定檢查點間隔。比如一個任務平均每 150 分鐘中斷一次,那麼每 15 分鐘保存一次會比較合理。這種做法既能減少重算損失,也不會帶來過多額外開銷。最終可行的間隔,還取決於你的儲存寫入速度。

如果我只保存模型權重,會發生什麼?

你會遺失關鍵的最佳化器狀態,包括 Adam 的動量值;學習率排程器會重設到初始階段;隨機數產生器會產生不同的資料打亂順序。這些缺失都會導致訓練結果不一致。因此,若想實現真正的恢復,必須始終保存完整的狀態字典。

我可以在不同數量的 GPU 上恢復訓練嗎?

可以,你可以在不同 GPU 數量的配置上恢復訓練。檢查點中保存的模型參數和最佳化器狀態可以跨配置遷移。不過,你需要相應調整 batch size 和學習率。資料載入器的分發方式可能會改變,但保存的 epoch 和模型狀態依然有效。

我怎麼判斷恢復是否成功?

先執行幾個 batch,並把損失值與當機前的日誌進行比較;梯度範數也應當盡量接近。同時觀察恢復初期是否出現 NaN 值或梯度爆炸。任何顯著偏差,都表示恢復流程存在錯誤,應立刻停止並排查,避免繼續浪費算力。

檢查點會拖慢訓練嗎?

檢查點機制確實會帶來一定開銷,但現代系統已經能把這部分成本壓得很低。傳統儲存方式通常會占用 5%–10% 的訓練時間,而高速 NVMe 分層儲存可將其降到 1% 以下。非同步檢查點還能把資料先卸載到主機記憶體,讓 GPU 在保存過程中繼續工作。