在 PyTorch 中視覺化並理解 GPU 記憶體 – Hugging Face Blog 摘要
TL;DR
本文展示了如何在 PyTorch 中記錄並視覺化 GPU 記憶體快照,解釋了記憶體分析中各個部分代表的意義(模型參數、優化器狀態、激活值、梯度、優化器中間變量),並提供了一個估算訓練期間峰值記憶體使用的通用公式。
🔎 PyTorch 視覺化工具
PyTorch 提供了一個內建工具來記錄並視覺化 GPU 記憶體使用情況。透過在執行程式碼前呼叫 torch.cuda.memory._record_memory_history 並在之後呼叫 torch.cuda.memory._dump_snapshot,你可以獲得一個 profile.pkl 檔案,並可以在 https://pytorch.org/memory_viz 查看。該視覺化工具會顯示記憶體分配與釋放的時間軸。
在一個簡單的 nn.Linear 層範例中,圖表顯示:
- 模型建立為權重和偏置分配了約 2 GB (float32)。
- 每個輸入張量增加約 200 MB。
- 每次前向傳播輸出增加約 1 GB。
- 來自先前步驟的激活值會保留到不再需要用於反向傳播為止,之後其記憶體會被釋放。
- 重新分配變數會釋放先前引用的張量。
📊 訓練期間的記憶體視覺化
對於使用大型語言模型 (Qwen/Qwen2.5-1.5B) 和 AdamW 優化器的實際訓練迴圈,記憶體分析顯示了三個峰值,每個訓練迭代出現一個。
記憶體分析可以分解如下:
- 模型初始化 – 模型參數 (藍色) 佔用記憶體並保持分配狀態直到訓練結束。
- 前向傳播 (Forward Pass) – 激活值 (橘色) 逐層計算並儲存;它們在損失計算時達到峰值。
- 反向傳播 (Backward Pass) – 計算梯度 (黃色);激活值被捨棄,導致橘色區域縮小。
- 優化器步驟 (Optimizer Step) – 優化器狀態 (綠色) 只初始化一次;優化器使用梯度來更新參數,並暫時儲存優化器中間變量 (紅色)。更新後,梯度和中間變量會被釋放。
這種模式在每次迭代中重複,產生了觀察到的峰值。
📐 估算記憶體需求
峰值記憶體使用量是分析中的最高點,根據 Batch Size 的不同,可能會發生在前向傳播或優化器步驟期間。
涵蓋這兩種情況的通用表達式為:
Total Memory = Model Memory + Optimizer State + max(Gradients + Optimizer Intermediates, Activations)
其中各個項目的定義如下。
模型參數
Model Memory = N × P
- N = 參數數量
- P = 位元組單位精度 (例如 float32 為 4)
以 Qwen2.5-1.5B 為例 (1.5B 參數,float32): Model Memory = 1.5 × 10⁹ × 4 bytes = 6 GB。
優化器狀態
對於 AdamW,它會為每個參數儲存兩個動量: Optimizer State Size = 2 × N × P
梯度
Gradients Memory = N × P (與模型參數大小相同)。
優化器中間變量
Optimizer Intermediates Memory = N × P (與模型參數大小相同)。
激活值
激活值記憶體取決於 Batch Size (B)、序列長度 (L) 和每個 Token 的激活值數量 (A)。 Activation Memory = A × B × L × P
A 可以透過 Forward Hooks 直接測量,但本文提供了一個從多個模型線性擬合得出的啟發式方法:
A = 4.6894 × 10⁻⁴ × N + 1.8494 × 10⁶
使用此啟發式方法,你可以在不執行完整前向傳播的情況下估算激活值記憶體。
總記憶體公式 (綜合)
代入各個組成公式後得到:
Total Memory = N×P + 2×N×P + max(N×P, N×P, A×B×L×P)
= 3×N×P + max(N×P, A×B×L×P)
由於 max 項至少會等於 N×P,因此當激活值佔主導地位時,表達式簡化為:
Total Memory = 3×N×P + A×B×L×P
否則,優化器步驟佔主導地位,總量為 4×N×P。
本文包含一個小工具,可以輸入 N, P, B, L 並獲得估算值。
🚀 下一步
理解記憶體分析能讓你尋找降低使用量的方法,例如降低 Batch Size、使用梯度檢查點 (Gradient Checkpointing) 或切換到混合精度訓練。TRL 文件中「減少記憶體使用量」章節中的建議廣泛適用於任何基於 PyTorch 的訓練。
🤝 致謝
感謝 Kashif Rasul 對本文提供的反饋與建議。