梯度檢查點
梯度檢查點(也稱為激活檢查點)是一種節省記憶體的技巧,它在前向傳播期間丟棄大多數中間激活,並在反向傳播期間動態重新計算它們。
概述
It lets you train deeper, larger networks by trading extra compute for much lower memory use.
深入探討
訓練神經網路通常會在前向傳播過程中儲存每一層的激活,因為反向傳播需要它們來計算梯度。對於深度模型來說,這些活化支配著記憶。相反,梯度檢查點僅在一組稀疏的“檢查點”層上保存激活,並丟棄其餘部分。當反向傳播到達啟動被刪除的區域時,它會重新執行該段的前向計算以重新產生所需的內容,然後繼續。大約每個 N 平方根層都放置檢查點,啟動記憶體從 N 階下降到 N 平方根階,而計算量僅增加約一次額外的前向傳遞(大約慢 20-30%)。這使得在同一 GPU 上適應更大的批量大小或更深的轉換器成為可能。
技術洞察
該技術利用了時間與記憶體的權衡。儲存所有啟動速度很快,但很耗記憶體;相對於耗盡記憶體的成本,在現代加速器上重新計算它們的成本很低。像 PyTorch (torch.utils.checkpoint) 這樣的框架包裝了一個模組,因此它的前向輸出被保存,但它的內部結構在後向過程中被重新計算。選擇檢查點放置很重要:大約 sqrt(N) 段的均勻間距可以最大限度地減少總內存,同時僅添加一個額外的前向計算整體。
戰略影響
成本與預算
多年來,架構決策決定著效能和營運成本。
更明確的決策
技術教育幫助團隊選擇正確的堆疊,而不僅僅是最新的堆疊。
品質管控
更好的工程選擇可以減少生產中的可靠性事故。
梯度檢查點的未來
梯度檢查點現在是大型模型訓練的標準,並且越來越自動化,函式庫會為您選擇最佳檢查點位置。它與 FSDP、混合精度和卸載自然配合,以提高模型尺寸。期望「選擇性」檢查點僅重新計算廉價的操作,同時快取昂貴的操作(如注意力矩陣),加上 PyTorch 的 torch.compile 等工具中的編譯器驅動方法,自動決定保存哪些內容與重新計算以獲得最佳速度記憶體平衡。
現實世界的實施
透過丟棄和重新計算層激活,在單一 GPU 上訓練具有更大批量大小的深度轉換器。
在高解析度影像上微調視覺模型,否則啟動圖會溢出 GPU 記憶體。
擁抱 Face Transformers 啟用gradient_checkpointing=True,以在微調期間適應十億參數模型。
將檢查點與 FSDP 結合,使參數和活化都保持較小,從而能夠訓練非常大的語言模型。
風險與防護欄
優化一項基準測試可以隱藏更廣泛的系統弱點。
基礎設施和維護成本常常被低估。
隨著系統變得更加複雜,安全性和可觀察性差距可能會擴大。
實施路線圖
在實施之前定義延遲、品質和成本目標。
在實際負載和資料條件下進行基準測試。
儀器監控錯誤、漂移和使用者影響。
在擴展之前準備回滾和事件回應路徑。
不斷探索
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 Gradient Checkpointing 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
常見問題
What is Gradient Checkpointing?
梯度檢查點(也稱為激活檢查點)是一種節省記憶體的技巧,它在前向傳播期間丟棄大多數中間激活,並在反向傳播期間動態重新計算它們。它可以讓您透過用額外的計算換取更低的記憶體使用來訓練更深、更大的網路。
為了節省內存,梯度檢查點主要交換什麼?
梯度檢查點在向後傳遞過程中重新計算丟棄的激活,花費額外的計算來換取顯著減少的記憶體。
為什麼激活通常在前向傳遞過程中儲存?
反向傳播使用前向傳遞的中間活化來計算梯度,因此它們必須可用,除非重新計算。
如果在 N 層網路中的每個 sqrt(N) 層放置檢查點,那麼啟動記憶體大致如何擴展?
每個 N 平方根層的間隔檢查點將儲存的啟動記憶體從 N 階減少到 sqrt(N) 階。
適當放置的梯度檢查點通常會增加大約多少額外計算?
如果檢查點位置良好,開銷大約是額外的前向傳遞,通常會減速 20-30% 左右。
在 PyTorch 中,哪個實用程式通常用於將梯度檢查點應用於模組?
torch.utils.checkpoint 包裝一個模組,以便在向後期間重新計算其內部啟動而不是儲存。