Accelerate ND-Parallel:高效多 GPU 訓練指南

Hugging Face 在 Accelerate 與 Axolotl 中引入了 ND-Parallelism,提供了一種簡化的方式,讓使用者能在單一訓練腳本中結合多種平行化策略——資料平行化 (DP)、完全分片資料平行化 (FSDP)、張量平行化 (TP) 與上下文平行化 (CP)——。此整合使開發者在跨多節點 GPU 叢集訓練擁有數十或數百億參數的模型時,能在記憶體使用與通訊開銷之間取得最佳平衡。

核心平行化策略

資料平行化 (DP)

資料平行化會在每個裝置上複製完整模型、梯度與優化器狀態。每個裝置處理不同的子批次資料,且在更新參數前會同步所有裝置的梯度。此方式提升吞吐量,但要求整個模型能容納於單一 GPU。

完全分片資料平行化 (FSDP)

FSDP 將模型權重、梯度與優化器狀態在 GPU 之間分片,降低每個裝置的記憶體占用。執行前向或反向傳播時,FSDP 會收集特定層(通常是 transformer 解碼器區塊)所需的權重,完成後再重新分片。此方式以較高的通訊開銷換取顯著降低的峰值記憶體使用量。

張量平行化 (TP)

張量平行化將大型線性層(例如前饋層或注意力投影)在裝置之間切分。與 FSDP 的動態分片不同,TP 產生靜態的記憶體分區。由於 TP 需要頻繁的激活同步,它在使用高頻寬連接(如 NVLink)的單一節點內最為有效,且不建議在 PCIe 連接的 GPU 上使用。

上下文平行化 (CP)

上下文平行化將輸入序列在 GPU 之間分片,以處理因注意力的二次方尺度而導致的極長序列,否則會超出 GPU 記憶體。透過 RingAttention,每個 GPU 持有 query、key、value 矩陣的一部分,並在 GPU 環形網路中循環傳遞 key‑value 分片,確保每個 query 能對整個序列計算注意力分數,同時分散計算與記憶體負載。

ND-Parallelism:組合策略以實現多節點擴展

多節點訓練常因節點間延遲與記憶體限制而成為瓶頸。ND-Parallelism 將叢集視為多維拓撲,以最佳化通訊。

混合分片資料平行化 (HSDP)

HSDP 是一種 2D 平行化方法,於單一節點內執行 FSDP(利用快速的節點內連接),並在節點間使用 DP。此方式將緩慢的節點間通訊縮減至單一梯度同步步驟,提升吞吐量,但相較於純 FSDP 會增加記憶體使用量。

FSDP + 張量平行化

結合 FSDP 與 TP 透過 FSDP 在節點間分片模型,並在單一節點內以 TP 切分層。此舉降低 FSDP 的延遲,允許訓練單一裝置無法容納的巨大層,且可減少全域批次大小。

FSDP + 上下文平行化

此 2D 策略在訓練極長序列時使用。雖然 CP 已與 FSDP 整合,於 CP 之上再加入 FSDP 可進一步降低模型權重與優化器狀態的記憶體需求。

混合分片資料平行化 + 張量平行化

此 3D 階層結構使用 DP 在節點群組間複製模型,FSDP 在群組內分片模型,TP 在每個節點內切分層。此配置提供了最大彈性,以因應特定硬體與擴展限制。

實作與使用說明

在 Accelerate 與 Axolotl 中的設定

使用者可以透過 Accelerate 中的 ParallelismConfig 類別或 Axolotl 中的特定設定欄位來配置這些策略:

  • dp_shard_size: Degree of FSDP
  • dp_replicate_size: Degree of DP
  • tp_size / tensor_parallel_size: Degree of TP
  • cp_size / context_parallel_size: Degree of CP

記憶體與穩定性最佳化

  • CPU RAM Efficient Loading: 對於過大無法放入單一裝置的模型,啟用 cpu_ram_efficient_loadingSHARDED_STATE_DICT(位於 FullyShardedDataParallelPlugin)是關鍵。
  • Effective Batch Size: 有效批次大小計算方式為 micro_batch_size * gradient_accumulation_steps * dp_world_size,其中 dp_world_size = (dp_shard_size * dp_replicate_size) / tp_size
  • Learning Rate Scaling: 隨著有效批次大小增大,學習率應以線性或平方根方式縮放,以維持穩定性。
  • Gradient Checkpointing: 此方式透過在反向傳播時重新計算中間激活,以計算換取記憶體。可將激活記憶體減少 60‑80%,但會使訓練時間增加約 20‑30%。

Sources