JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

別再複製本可分片的東西:ZeRO 與 FSDP

純資料平行因為把每樣東西都存了 N 份相同副本而浪費記憶體。改成把優化器、梯度與權重分片,並在需要的瞬間才取回。

資料平行白白浪費的冗餘

在純資料平行裡,N 份副本中的每一份都持有完整的優化器狀態、完整的梯度與完整的權重。在 all-reduce 之後,每張裝置都擁有逐位元相同的梯度——所以其中 *N − 1* 份純屬浪費。零冗餘優化器(ZeRO) 背後的洞見是:資料平行並不要求複製,它只要求在某個張量被需要的當下,正確的裝置擁有它。其餘一切都能在 N 個 rank 之間分片,按需求再取回。

M_{\text{replica}} = \underbrace{2\Psi}_{\text{fp16 weights}} + \underbrace{2\Psi}_{\text{fp16 grads}} + \underbrace{12\Psi}_{\text{fp32 master, }m,v} = 16\Psi

混合精度 Adam 讓每個資料平行副本為每個參數儲存 16 位元組——這正是 ZeRO 要消除的冗餘。

ZeRO 的三個階段

  1. 階段 1——分片優化器狀態。每個 rank 只擁有 1/N 參數的 Adam 動量,也只更新那一片。光是這點就消去了四項記憶體中最大的一項。
  2. 階段 2——連梯度也分片。用 reduce-scatter 取代全域 all-reduce,讓每個 rank 只收到自己負責更新的那一片梯度。
  3. 階段 3——連參數也分片。沒有任何 rank 持有完整模型;某層的權重在該層執行前才 all-gather 取回,執行後立刻釋放。
\begin{aligned} P_{\text{os}}\ (\text{Stage 1}) &: 4\Psi + \frac{12\Psi}{N}\\ P_{\text{os}+g}\ (\text{Stage 2}) &: 2\Psi + \frac{14\Psi}{N}\\ P_{\text{os}+g+p}\ (\text{Stage 3}) &: \frac{16\Psi}{N} \end{aligned}

三個 ZeRO 階段下的每裝置記憶體:每個階段多切分一項,階段 3 把所有項都除以 N。

階段 3 才是記憶體真正崩塌之處:每張裝置的峰值權重記憶體下降 N 倍,於是增加裝置能讓你訓練更大的模型,而不只是更大的批次。代價是更多通訊——前向一次 all-gather、反向再一次——這正是為什麼上一篇的重疊紀律在這裡沒有商量餘地。廣為使用的實作是 DeepSpeed-ZeRO,它還提供 CPU 與 NVMe 卸載,在你願意用頻寬換容量時,能把分片狀態完全推離 GPU。

FSDP:同樣的想法,以層為單位

完全分片資料平行(FSDP) 就是把 ZeRO 階段 3 的想法,表達成對模型各模組的包裝。你把層分組成 FSDP 單元;每個單元的參數分片存放在各 rank 上,在該單元的前向(或反向)執行前才 all-gather 成完整張量,之後立即重新分片,使任一時刻只有一個單元的完整權重駐留。由於單元邊界同時決定了記憶體與 all-gather 的粒度,選擇如何包裝便是主要的效能旋鈕:太細會付出啟動開銷,太粗則暫態的完整權重尖峰會變大。

\underbrace{2\Psi}_{\text{DP all-reduce}} \;\longrightarrow\; \underbrace{\Psi}_{\text{all-gather (fwd)}} + \underbrace{\Psi}_{\text{all-gather (bwd)}} + \underbrace{\Psi}_{\text{reduce-scatter}} = \underbrace{3\Psi}_{1.5\times}

FSDP/ZeRO-3 用 3Ψ 的全聚合 + 歸約-散射模式取代普通資料平行的 2Ψ 全歸約——通訊量為 1.5 倍。

# Wrap each transformer block as its own FSDP unit
for block in model.transformer_blocks:
    fsdp_wrap(block)               # params now sharded across ranks
fsdp_wrap(model)                   # root unit handles embeddings/head

# Forward of one block, conceptually:
#   all_gather(block.params) -> compute -> reshard(block.params)
以區塊為單位的 FSDP 包裝:任一時刻只實體化一個區塊的完整權重。

兩個與分片搭配使用的記憶體槓桿

分片對付的是權重/梯度/優化器這幾項,但對激活值毫無幫助——那要靠 激活值重算,我們會在模型平行那篇完整展開。兩者互補:ZeRO/FSDP 縮小參數的足跡,而重算縮小激活值的足跡,大型訓練會同時用上兩者。分片也重塑了你如何把狀態存到磁碟:沒有任何單一 rank 擁有完整模型,因此檢查點以 分片檢查點 的形式寫出,每個 rank 平行傾印自己的那一片。

\underbrace{O(L)}_{\text{store every layer}} \;\longrightarrow\; \underbrace{O(\sqrt{L})}_{\sqrt{L}\text{ checkpoints}} \quad (\text{one extra forward pass})

激活重計算只保留 √L 個檢查點,以一次額外前向傳播將激活記憶體從 O(L) 降到 O(√L)。