使用 PyTorch Fully Sharded Data Parallel 加速大型模型訓練
Hugging Face 已將 PyTorch Fully Sharded Data Parallel (FSDP) 整合至 Accelerate 函式庫中,透過將優化器狀態、梯度和參數分片(sharding)到資料並行工作節點上,讓開發者能夠訓練顯著更大的模型。此整合透過啟用更大的批次大小(batch size),以及透過 CPU offloading 能力訓練原本會超出 GPU 記憶體限制的模型,使大型模型訓練變得更加普及。
FSDP 與 Distributed Data Parallel (DDP) 的比較
PyTorch FSDP 透過消除每個 GPU 上完整模型副本所產生的冗餘記憶體消耗,改進了 Distributed Data Parallel (DDP)。
在 DDP 中,每個工作節點都維護模型參數、梯度和優化器狀態的完整副本。雖然每個工作節點處理不同的資料批次,但在更新模型之前,它們必須執行 all-reduce 操作來平均所有工作節點的梯度。
在 FSDP 中,優化器狀態、梯度和模型參數會在工作節點之間進行分片。在正向和反向傳播期間,FSDP 使用 all-gather 操作僅檢索特定層或封裝模組所需的參數,並在計算後立即釋放它們。接著,透過 reduce-scatter 操作對局部梯度進行平均並分發,使每個工作節點僅更新其局部的參數分片。這種方法大幅降低了每個 GPU 的記憶體佔用。
GPT-2 的效能基準測試
Hugging Face 使用兩張 NVIDIA Titan RTX GPU(各 24GB)在因果語言模型(Causal Language Modeling)任務上,針對 FSDP 與 DDP 進行了基準測試。
GPT-2 Large (762M Parameters)
FSDP 與 DDP 相比,能實現顯著更大的批次大小。在不使用 CPU offload 的情況下,FSDP 允許的批次大小最高可達 15(DDP 為 7)。啟用 CPU offload 後,批次大小進一步增加到 22。雖然混合精度(FP16)下的 DDP 在原始訓練時間上最快,但 FSDP 提供了大型批次所需的記憶體效率,這對於具有動態批次(dynamic batching)的應用特別有益。
GPT-2 XL (1.5B Parameters)
對於 GPT-2 XL 模型,DDP 即使在批次大小為 1 時也會發生 CUDA Out of Memory (OOM) 錯誤。相比之下,FSDP 實現了成功的訓練:
- FSDP (Zero-Stage 3): 在 2 張 GPU 上支援每張 GPU 批次大小為 5。
- FSDP with CPU Offload: 在單張 GPU 上支援批次大小為 10 的訓練,以及在 2 張 GPU 上每張 GPU 批次大小為 14 的訓練。
技術實作與配置
透過 Accelerate 進行整合
使用者可以透過 accelerate config CLI 或透過 FullyShardedDataParallelPlugin 進行更細粒度的控制。關鍵配置選項包括:
- Sharding Strategy: 在
FULL_SHARD與SHARD_GRAD_OP之間進行選擇。 - Min Num Params: 讓層被預設的自動封裝策略(auto-wrap policy)封裝所需的最小參數數量。
- Offload Params: 一個布林值,決定是否應將參數和梯度卸載(offload)到 CPU。
自動封裝策略的角色
min_num_params 設定對於記憶體優化至關重要。當使用 default_auto_wrap_policy 時,FSDP 會在參數數量超過指定閾值時封裝該層。在 BERT-Large (330M) 的基準測試顯示,使用 auto-wrap 的 FSDP 消耗的記憶體約為 DDP 的一半。降低 min_num_params(例如,降至 2k)與較高的閾值(例如,1M)相比,可以略微進一步降低記憶體使用量。
關鍵注意事項與限制
使用 FSDP 的開發者必須注意以下幾項技術限制:
- Optimizer Initialization: FSDP 會將參數扁平化並進行原地分片。因此,模型必須在建立優化器 之前 透過
accelerator.prepare(model)進行準備。在封裝模型之前建立優化器可能會導致優化器失效或導致記憶體使用量增加。 - Parameter Groups: 由於 FSDP 會將嵌套模組扁平化為 1D 陣列,在封裝之前建立的參數組(例如,對偏置項應用不同的權重衰減)會被合併成單一組別並遺失。
- Multiple Models: 訓練多個模型時,在建立各自的優化器之前,必須先準備模型,以避免錯誤。
- Mixed Precision: 在本文發布時,由於 PyTorch 尚待修復,FSDP 目前尚不支持混合精度。
分布式訓練方法的總結
FSDP 是更廣泛的分布式訓練策略生態系統的一部分,旨在處理巨型模型:
- ZeRO (Zero Redundancy Optimizer): FSDP 的基礎,對優化器狀態(Stage 1)、梯度(Stage 2)和參數(Stage 3)進行分片。
- Tensor Parallelism: 將個別大型層的參數分片到不同 GPU 上。
- Pipeline Parallelism: 將不同的層分佈在不同的 GPU 上,並透過流水線方式處理微批次(micro-batches)。
- 3D Parallelism: ZeRO、Tensor 與 Pipeline 並行技術的結合,用於處理擁有數千億參數的模型。
Sources
相關
- Dispatch
- Dispatch
- Dispatch
- Dispatch
- Dispatch