使用 PyTorch FSDP 微調 Llama 2 70B
Hugging Face 詳細介紹了一種使用 PyTorch Fully Sharded Data Parallelism (FSDP) 微調 Llama 2 70B 模型的方法,並利用了 Hugging Face Transformers、Accelerate 和 TRL 函式庫。這種方法透過將優化器狀態 (optimizer states)、梯度 (gradients) 和參數 (parameters) 分散到各個裝置上,實現了在多節點、多 GPU 設定下訓練大規模模型的能力。
克服模型載入期間的 CPU RAM 瓶頸
載入 Llama 2 70B 模型通常需要大量的 CPU RAM;如果節點上的每個進程都載入模型,可能需要大約 2TB 的 CPU RAM (70B parameters * 4 bytes * 8 GPUs)。為了防止記憶體不足 (OOM) 錯誤,Hugging Face 在 transformers 和 accelerate 中實作了一種特定的初始化策略:
- Meta Device 初始化:在所有 rank 上使用
metadevice 建立模型,這意味著初始化時不包含權重。 - Rank 0 載入:僅在 rank 0 上載入 state dict。
- 空參數配置:所有其他 rank 使用
torch.empty()在metadevice 上建立空參數。 - 狀態廣播:透過設置
sync_module_states=True,FSDP 會在訓練開始前將權重從 rank 0 廣播到所有其他 rank。
這種方法確保每個節點只有一個進程將預訓練模型載入到 CPU RAM 中,大幅減少了設定階段的記憶體佔用。
使用分片 State Dicts 進行高效檢查點儲存
使用 FULL_STATE_DICT 並在 rank 0 上進行 CPU offloading 來儲存完整的中間檢查點 (checkpoints),往往會導致 NCCL Timeout 錯誤和嚴重的延遲。為了解決這個問題,Hugging Face 建議在 FSDP 配置中使用 SHARDED_STATE_DICT。
- 中間檢查點:
SHARDED_STATE_DICT會分別儲存每個 GPU 的分片,從而實現更快的訓練儲存與恢復。 - 最終模型匯出:為了獲得用於部署的標準模型 state dict,僅在訓練結束並呼叫
trainer.save_model()之前,才將 state dict 類型切換為FULL_STATE_DICT。
優化 VRAM 與訓練速度
為了降低計算成本並提高訓練速度,此實作採用了兩種主要技術:梯度檢查點 (Gradient Checkpointing) 和 Flash Attention。
Flash Attention
標準的注意力機制由於在逐元素操作(masking、softmax 和 dropout)期間存在冗餘的高頻寬記憶體 (HBM) 讀寫,通常受限於記憶體頻寬。Flash Attention 透過以下方式進行優化:
- Kernel Fusion:它將中間步驟保留在 SRAM 中,並且僅將最終結果寫回 HBM 一次。
- Tiling:將 NxN 的 softmax/scores 計算切分成區塊以符合 SRAM 限制,並利用線上 softmax (online softmax) 演算法。
- Recomputation:反向傳播會重新計算必要的數值,而不是儲存來自前向傳播的整個 NxN 矩陣,從而顯著減少記憶體消耗。
Gradient Checkpointing
啟用梯度檢查點可以進一步減少 VRAM 使用量,在微調 70B 參數模型期間允許更大的 batch size 或更長的序列長度。
實作與硬體規格
硬體配置
微調是使用以下硬體進行的:
節點:2 個節點(至少需要 1 個)。
GPU:每個節點 8 個 A100 (80GB) GPU。
互連:NVLink (節點內) 和 Elastic Fabric Adapter (節點間)。
系統 RAM:每個節點 1TB。
CPU:每個節點 96 核心。
訓練執行
訓練過程使用了 accelerate launch 指令,並採用 FULL_SHARD 策略和 TRANSFORMER_BASED_WRAP 自動封裝策略。啟用了 bf16 混合精度訓練。對於具有 8 個 A100 80GB GPU 的單節點設定,建議使用來自 bitsandbytes 的 paged_adamw_32bit 優化器來管理記憶體。
使用 meta-llama/Llama-2-70b-chat-hf 模型和 smangrul/code-chat-assistant-v1 資料集,大約在 13.5 小時內完成了微調。