資料平行白白浪費的冗餘
在純資料平行裡,N 份副本中的每一份都持有完整的優化器狀態、完整的梯度與完整的權重。在 all-reduce 之後,每張裝置都擁有逐位元相同的梯度——所以其中 *N − 1* 份純屬浪費。零冗餘優化器(ZeRO) 背後的洞見是:資料平行並不要求複製,它只要求在某個張量被需要的當下,正確的裝置擁有它。其餘一切都能在 N 個 rank 之間分片,按需求再取回。
混合精度 Adam 让每个数据并行副本为每个参数存储 16 字节——这正是 ZeRO 要消除的冗余。
ZeRO 的三個階段
- 階段 1——分片優化器狀態。每個 rank 只擁有 1/N 參數的 Adam 動量,也只更新那一片。光是這點就消去了四項記憶體中最大的一項。
- 階段 2——連梯度也分片。用 reduce-scatter 取代全域 all-reduce,讓每個 rank 只收到自己負責更新的那一片梯度。
- 階段 3——連參數也分片。沒有任何 rank 持有完整模型;某層的權重在該層執行前才 all-gather 取回,執行後立刻釋放。
三个 ZeRO 阶段下的每设备内存:每个阶段多切分一项,阶段 3 把所有项都除以 N。
階段 3 才是記憶體真正崩塌之處:每張裝置的峰值權重記憶體下降 N 倍,於是增加裝置能讓你訓練更大的模型,而不只是更大的批次。代價是更多通訊——前向一次 all-gather、反向再一次——這正是為什麼上一篇的重疊紀律在這裡沒有商量餘地。廣為使用的實作是 DeepSpeed-ZeRO,它還提供 CPU 與 NVMe 卸載,在你願意用頻寬換容量時,能把分片狀態完全推離 GPU。
FSDP:同樣的想法,以層為單位
完全分片資料平行(FSDP) 就是把 ZeRO 階段 3 的想法,表達成對模型各模組的包裝。你把層分組成 FSDP 單元;每個單元的參數分片存放在各 rank 上,在該單元的前向(或反向)執行前才 all-gather 成完整張量,之後立即重新分片,使任一時刻只有一個單元的完整權重駐留。由於單元邊界同時決定了記憶體與 all-gather 的粒度,選擇如何包裝便是主要的效能旋鈕:太細會付出啟動開銷,太粗則暫態的完整權重尖峰會變大。
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)兩個與分片搭配使用的記憶體槓桿
分片對付的是權重/梯度/優化器這幾項,但對激活值毫無幫助——那要靠 激活值重算,我們會在模型平行那篇完整展開。兩者互補:ZeRO/FSDP 縮小參數的足跡,而重算縮小激活值的足跡,大型訓練會同時用上兩者。分片也重塑了你如何把狀態存到磁碟:沒有任何單一 rank 擁有完整模型,因此檢查點以 分片檢查點 的形式寫出,每個 rank 平行傾印自己的那一片。
激活重计算只保留 √L 个检查点,以一次额外前向传播将激活内存从 O(L) 降到 O(√L)。