使用 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 在 transformersaccelerate 中實作了一種特定的初始化策略:

  1. Meta Device 初始化:在所有 rank 上使用 meta device 建立模型,這意味著初始化時不包含權重。
  2. Rank 0 載入:僅在 rank 0 上載入 state dict。
  3. 空參數配置:所有其他 rank 使用 torch.empty()meta device 上建立空參數。
  4. 狀態廣播:透過設置 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 的單節點設定,建議使用來自 bitsandbytespaged_adamw_32bit 優化器來管理記憶體。

使用 meta-llama/Llama-2-70b-chat-hf 模型和 smangrul/code-chat-assistant-v1 資料集,大約在 13.5 小時內完成了微調。

Sources