ZeRO-1/2/3(Zero Redundancy Optimizer 三階段)
3 秒看懂
ZeRO 是微軟 DeepSpeed 團隊提出的分散式資料並行記憶體最佳化技術(OSDI 2020),核心思想極簡:標準資料並行中每個 GPU 都持有完整的模型副本(最佳化器狀態+梯度+引數),ZeRO 把這些狀態分片到各 GPU 上,按需 AllGather 拼回。三個階段逐步加大分片範圍,記憶體從”每 GPU 持有全量”下降到近似”全量/N”。
| 一句話 | 分片內容 | 記憶體節省 |
|---|---|---|
| Stage 1 | 最佳化器狀態 | 約 4× |
| Stage 2 | + 梯度 | 約 8× |
| Stage 3 | + 引數 | 理論線性(×N) |
N = 資料並行度(參與 ZeRO 分片的 GPU 數量);具體倍數取決於最佳化器型別與精度配置,此處以混合精度 + Adam 為例的典型估算。
3 分鐘產業解釋
為什麼 ZeRO 存在?
大型模型訓練的核心矛盾:單卡裝不下,模型並行又太複雜。
傳統的分散式訓練方案主要有兩條路:
- 模型並行(Tensor/Pipeline Parallelism):把模型切開放到不同 GPU。精度無損,但需要大量侵入式程式碼改造,通訊拓撲設計複雜。
- 資料並行(Data Parallelism, DP):每個 GPU 持有完整模型,各算各的梯度,再 AllReduce 同步。實現簡單,但每個 GPU 的記憶體冗餘極大。
以混合精度訓練 + Adam 最佳化器為例,每個引數需要:
| 元件 | 精度 | 每引數位元組數 |
|---|---|---|
| 引數(fp16) | 16-bit | 2 B |
| 梯度(fp16) | 16-bit | 2 B |
| Adam 一階動量(fp32) | 32-bit | 4 B |
| Adam 二階動量(fp32) | 32-bit | 4 B |
| fp32 主權重副本 | 32-bit | 4 B |
| 合計 | 16 B |
此處 16B/param 是混合精度(fp16 引數/梯度 + fp32 最佳化器狀態)+ Adam 的經典拆解,見 ZeRO 原始論文(Rajbhandari et al., OSDI 2020)。如果使用 SGD+Momentum,最佳化器狀態部分減半。
一個 175B 引數的模型(GPT-3 量級)→ 每個 GPU 需要 175B × 16B = 2,800 GB,遠超任何單 GPU 視訊記憶體。即使只考慮 7B 引數,也需要 112 GB,超過單張 80GB A100 的容量。
ZeRO 的回答是:既然 DP 中每個 GPU 都存了完整的 16Φ 位元組,何不把它切分成 N 份,每個 GPU 只存 16Φ/N?需要用時再 AllGather 回來。
產業價值:ZeRO 讓資料並行也能訓練遠超單卡視訊記憶體的模型,且不需要侵入式修改模型程式碼。它與模型並行(Tensor/Pipeline Parallelism)可以疊加使用,是當前大型模型訓練基礎設施的關鍵拼圖之一。
15 分鐘專家深入
三階段機制詳解
設模型引數量為 Φ,資料並行度為 N_d(參與 ZeRO 分片的 GPU 數),混合精度 + Adam 訓練。
Stage 1:分片最佳化器狀態(ZeRO-OS)
分片物件:Adam 的一階動量、二階動量、fp32 主權重副本(共 12Φ 位元組)。
機制:
- 每個 GPU 只儲存 1/N_d 的最佳化器狀態(對應 1/N_d 的引數切片)
- 前向/反向傳播時,每個 GPU 使用本地持有的 fp16 引數(完整引數仍在每張卡上)
- 反向傳播完成後,每張卡得到完整的本地梯度(fp16)
- ReduceScatter:各卡將梯度按引數切片做規約,每個 GPU 拿到屬於自己負責的那 1/N_d 引數切片的聚合梯度
- 用聚合後的梯度更新本地負責的那部分 fp32 最佳化器狀態和 fp32 主權重
- 將更新後的 fp32 主權重截斷為 fp16,寫回本地 fp16 引數對應切片
- AllGather:各卡對更新後的 fp16 引數切片執行 AllGather,所有 GPU 獲得完整的更新後 fp16 引數
記憶體:
Stage 1 per-GPU memory = 2Φ + 2Φ + 12Φ/N_d = 4Φ + 12Φ/N_d
(params) (grads) (opt states)
通訊開銷:與標準 DP 的 AllReduce 相同(AllReduce = ReduceScatter + AllGather,總通訊量 2Φ/GPU)。Stage 1 沒有額外通訊。[來源:ZeRO 原始論文]
Stage 2:分片最佳化器狀態 + 梯度(ZeRO-OS+G)
在 Stage 1 基礎上進一步分片:梯度也按引數切片分割槽,每個 GPU 只儲存屬於自己的 1/N_d 引數切片對應的梯度。
機制:
- 前向傳播:每張卡持有完整 fp16 引數,正常計算
- 反向傳播:每張卡計算本地梯度(完整 Φ 個引數的梯度)
- ReduceScatter:各卡梯度按切片規約,每個 GPU 只保留自己負責的 1/N_d 切片的聚合梯度,其餘梯度立即釋放
- 用聚合梯度更新本地最佳化器狀態
- 更新 fp16 引數
記憶體:
Stage 2 per-GPU memory = 2Φ + (2Φ + 12Φ)/N_d = 2Φ + 14Φ/N_d
(params) (grads+opt states)
通訊開銷:仍與標準 DP AllReduce 等價。ReduceScatter 聚合梯度(Φ)+ 後續引數更新後隱式同步(Φ),總 2Φ/GPU。[來源:ZeRO 原始論文]
關鍵洞察:Stage 2 的核心優勢在於釋放了 (N_d-1)/N_d 的梯度記憶體。對於大型模型,梯度記憶體佔比與引數記憶體相同(都是 2Φ),這一步節省非常顯著。
Stage 3:分片一切(ZeRO-OS+G+P)
終極形態:引數也按切片分片,每個 GPU 只儲存 1/N_d 的 fp16 引數。
機制:
- 前向傳播前:AllGather 拼回當前層(或整個模型)的完整 fp16 引數 → 計算 → 釋放非本切片引數
- 反向傳播前:再次 AllGather 拼回完整引數 → 計算梯度 → 釋放非本切片引數
- ReduceScatter:梯度按切片規約 → 每卡只保留 1/N_d 聚合梯度
- 更新本地最佳化器狀態和引數
記憶體:
Stage 3 per-GPU memory = (2Φ + 2Φ + 12Φ)/N_d = 16Φ/N_d
(所有狀態均分片)
通訊開銷:高於標準 DP。每張卡每步額外需要兩次 AllGather(前向 + 反向各一次)來拼回引數:
Stage 3 通訊量 = AllGather_fwd(Φ) + AllGather_bwd(Φ) + ReduceScatter(Φ) = 3Φ/GPU
標準 DP 通訊量 = AllReduce(Φ) = ReduceScatter(Φ) + AllGather(Φ) = 2Φ/GPU
通訊比 = 3Φ / 2Φ = 1.5×
[來源:ZeRO 原始論文,§4.3]
三階段記憶體公式彙總
設 N_d 為 ZeRO 並行度,混合精度 + Adam(16B/param):
| 階段 | 每 GPU 記憶體 | N_d=64 時佔比 |
|---|---|---|
| 無 ZeRO | 16Φ | 100% |
| Stage 1 | 4Φ + 12Φ/N_d | ~26.2% |
| Stage 2 | 2Φ + 14Φ/N_d | ~13.9% |
| Stage 3 | 16Φ/N_d | ~1.6% |
N_d 越大,Stage 3 的記憶體優勢越極端。但 Stage 3 也引入了 1.5× 通訊量和更頻繁的 AllGather,實際吞吐受通訊頻寬制約。
ZeRO-Offload 與 ZeRO-Infinity
ZeRO 架構的後續擴充套件(同團隊):
- ZeRO-Offload(2021):將 fp32 最佳化器狀態和部分計算 offload 到 CPU 記憶體和 CPU 計算,使單 GPU 也能訓練大型模型。代價是 CPU-GPU PCIe 頻寬瓶頸。
- ZeRO-Infinity(2021):進一步擴充套件到 NVMe SSD,實現”無限”模型規模(論文標題字面含義)。利用 GPU 視訊記憶體 + CPU 記憶體 + NVMe 的層次化儲存。
ZeRO++
微軟後續工作,主要最佳化 Stage 3 的通訊效率:
- qwZ(Quantized Weights ZeRO):對 AllGather 傳輸的權重做量化(如 int8/int4),減少通訊量
- hpZ(Hierarchical Partitioning ZeRO):在節點內和節點間做兩級分片,優先在節點內(NVLink 高頻寬)做 AllGather,減少跨節點通訊
[注:ZeRO++ 的具體量化精度和效能資料,以 DeepSpeed 官方 GitHub 為準]
技術原理(最深)
資料並行中的記憶體冗餘問題
標準資料並行(DP)的工作流:
┌─────────────────────────────────────────────────┐
│ 標準 Data Parallel │
│ │
│ GPU-0 GPU-1 ... GPU-(N-1) │
│ ┌──────┐ ┌──────┐ ┌──────┐ │
│ │完整Φ │ │完整Φ │ │完整Φ │ │
│ │完整G │ │完整G │ │完整G │ │
│ │完整OS│ │完整OS│ │完整OS│ │
│ └──────┘ └──────┘ └──────┘ │
│ ↓ ↓ ↓ │
│ AllReduce(g) AllReduce(g) AllReduce(g) │
│ ↓ ↓ ↓ │
│ 本地更新 本地更新 本地更新 │
│ │
│ Φ=引數, G=梯度, OS=最佳化器狀態 │
│ 冗餘: 每張卡存 16Φ, N張卡共存 16Φ×N │
└─────────────────────────────────────────────────┘
ZeRO 的核心洞察:AllReduce 之後每張卡都有相同的聚合梯度,用於更新相同的引數和最佳化器狀態 → N-1 份是完全冗餘的。
ZeRO Stage 3 的 AllGather 與釋放機制
以一個簡化的兩層模型(Layer1, Layer2)為例,Stage 3 中每張卡只存 1/N 的引數分片:
Forward Pass (以 GPU-0 的視角):
┌─ AllGather Layer1 params ──→ 得到完整 L1 引數 ─┐
│ │
│ Compute Layer1 forward │
│ │
├─ 釋放非本切片的 L1 引數分片 │
│ │
├─ AllGather Layer2 params ──→ 得到完整 L2 引數 ─┐ │
│ │ │
│ Compute Layer2 forward │ │
│ │ │
└─ 釋放非本切片的 L2 引數分片 ┘ ┘
Backward Pass (類似):
AllGather Layer2 params → Compute L2 backward → 釋放
AllGather Layer1 params → Compute L1 backward → 釋放
ReduceScatter gradients → 保留 1/N_d 聚合梯度
注意:在實際實現(DeepSpeed)中,AllGather 可以按層粒度或按模型整體粒度執行。按層粒度可以降低峰值記憶體(不需要同時拼回所有層的引數),但增加通訊輪次。[具體粒度實現細節見 DeepSpeed 原始碼]
AllReduce vs ReduceScatter + AllGather 的通訊等價性
標準 AllReduce = ReduceScatter + AllGather
ReduceScatter: 每個 GPU 貢獻自己的資料,收到一個切片的聚合結果
通訊量: Φ/GPU(在環形拓撲下)
AllGather: 每個 GPU 貢獻一個切片,收到完整的聚合資料
通訊量: Φ/GPU(在環形拓撲下)
總計: 2Φ/GPU
ZeRO-1/2 的 ReduceScatter 不是"額外"的——它替代了 AllReduce 的前半段,
而後半段 AllGather 在引數更新後以不同形式發生。
總通訊量不變,仍為 2Φ/GPU。
通訊量對比(Stage 3 vs 基線 DP)
標準 DP:
Forward: 無通訊
Backward: AllReduce = ReduceScatter(Φ) + AllGather(Φ) = 2Φ/GPU
總計: 2Φ/GPU
ZeRO Stage 3:
Forward: AllGather(Φ) ← 拼回完整引數
Backward: AllGather(Φ) ← 再次拼回(因前向後已釋放)
+ ReduceScatter(Φ) ← 聚合梯度
總計: 3Φ/GPU
通訊比: 3Φ / 2Φ = 1.5×
技術演進史
| 時間 | 里程碑 | 關鍵進展 |
|---|---|---|
| 2020 | ZeRO 論文發表 (OSDI 2020) | 提出 Stage 1/2/3,在 1000 GPU 上訓練 1T 引數模型(記憶體層面) |
| 2020 | DeepSpeed v0.1-v0.3 | 開源實現 ZeRO Stage 1/2/3 |
| 2021.02 | ZeRO-Offload | 將最佳化器狀態 offload 到 CPU,單 GPU 訓練 10× 大型模型 |
| 2021.04 | ZeRO-Infinity | 擴充套件到 NVMe,突破 GPU+CPU 記憶體瓶頸 |
| 2021-2022 | DeepSpeed 與 Megatron-LM 融合 | Megatron-DeepSpeed:ZeRO + 張量並行 + 流水線並行 + 序列並行,成為萬億引數訓練主流方案 |
| 2022 | PyTorch FSDP | Meta 實現了 ZeRO-3 等效的 Fully Sharded Data Parallel,整合入 PyTorch 原生架構 |
| 2023 | ZeRO++ | 量化權重通訊 + 層次化分片,最佳化 Stage 3 通訊效率 |
| 2023+ | DeepSpeed-Chat / DeepSpeed-FastGen | ZeRO 在 RLHF 和推論場景的擴充套件應用 |
Megatron-DeepSpeed 被用於訓練 BLOOM-176B(BigScience 專案),這是 ZeRO + 模型並行組合在公開大型模型訓練中的標誌性案例。
技術路線對比(量化表)
ZeRO 三階段橫向對比
| 維度 | Stage 1 | Stage 2 | Stage 3 |
|---|---|---|---|
| 分片內容 | 最佳化器狀態 | 最佳化器狀態 + 梯度 | 最佳化器狀態 + 梯度 + 引數 |
| 每 GPU 記憶體 (混合精度+Adam) | 4Φ + 12Φ/N_d | 2Φ + 14Φ/N_d | 16Φ/N_d |
| 通訊量/GPU/step | 2Φ(同基線) | 2Φ(同基線) | 3Φ(1.5× 基線) |
| 適用場景 | 最佳化器狀態佔比大、引數記憶體還裝得下 | 大型模型、引數記憶體緊張 | 超大型模型、需極致節省記憶體 |
| 程式碼侵入性 | 低 | 低 | 低(但需要模型能被分片) |
| 典型 N_d 值 | 8-64 | 64-256 | 256+ 或跨節點 |
ZeRO vs 其他並行策略
| 策略 | 記憶體效率 | 通訊開銷 | 程式碼改造 | 精度影響 | 典型使用 |
|---|---|---|---|---|---|
| 標準 DP | 低(全量冗餘) | 2Φ/GPU | 無 | 無 | 小模型 |
| ZeRO-3 | 極高 | 3Φ/GPU(1.5×) | 極小 | 無 | 大型模型資料並行層 |
| Tensor Parallel | 中等 | 集合通訊每層 | 高(改模型結構) | 無 | 節點內層內並行 |
| Pipeline Parallel | 中等 | 點對點每 micro-batch | 中等 | 可能有 bubble 效率損失 | 跨節點層間並行 |
| FSDP (PyTorch) | 等效 ZeRO-3 | 等效 ZeRO-3 | 中等(PyTorch API) | 無 | PyTorch 生態使用者 |
| 專家並行 (MoE) | 啟用引數少 | All-to-All dispatch | 高 | 無 | MoE 架構模型 |
注意:上表中 ZeRO 和 Tensor/Pipeline Parallel 不是互斥的,實際大規模訓練中通常組合使用。常見組合:Tensor Parallel(節點內 NVLink)+ Pipeline Parallel(跨節點)+ ZeRO(資料並行維度)。
上下游
上游(ZeRO 依賴什麼)
| 層級 | 元件 | 說明 |
|---|---|---|
| 硬體 | GPU 視訊記憶體 | ZeRO 本質是視訊記憶體最佳化,視訊記憶體越大,可訓練模型越大 |
| 硬體 | GPU 間互聯頻寬 | NVLink(節點內)、InfiniBand/RoCE(節點間)直接影響 Stage 3 的實際吞吐 |
| 架構 | 深度學習架構 | PyTorch(DeepSpeed 基於此)、JAX/XLA 等 |
| 最佳化器 | Adam/AdamW | 最佳化器狀態是 ZeRO 分片的主要物件;SGD+Momentum 狀態更小,ZeRO 收益相對低 |
| 精度 | 混合精度訓練 | fp16/bf16 引數+fp32 最佳化器狀態的混合精度方案使最佳化器狀態佔比最大(12/16 = 75%) |
下游(ZeRO 支撐什麼)
| 應用 | 說明 |
|---|---|
| 大語言模型預訓練 | GPT-3/BLOOM/LLaMA 量級的訓練 |
| 大型模型微調(SFT/RLHF) | 全引數微調需要等量記憶體,ZeRO 可節省 |
| 長序列訓練 | 序列長度增加→啟用記憶體增加→引數/梯度/最佳化器記憶體更需節省 |
| MoE 訓練 | MoE 模型總引數量極大,但每個 token 只啟用部分專家;ZeRO 可處理專家引數的分片 |
關鍵指標
記憶體指標
| 指標 | 定義 | ZeRO 最佳化方向 |
|---|---|---|
| 每 GPU 模型狀態記憶體 | 引數 + 梯度 + 最佳化器狀態 | 隨 Stage 升級遞減 |
| 每 GPU 啟用記憶體 | 前向傳播中間結果(用於反向) | ZeRO 不直接最佳化(可用 activation checkpointing 配合) |
| 峰值記憶體 | 模型狀態 + 啟用 + 臨時緩衝 | ZeRO 降低模型狀態部分 |
通訊指標
| 指標 | 定義 | ZeRO 的影響 |
|---|---|---|
| 通訊量/GPU/step | 每步每 GPU 需傳送/接收的資料總量 | Stage 1/2 不增加,Stage 3 增加 1.5× |
| 通訊/計算比 | 通訊時間 / 計算時間 | 取決於模型大小和 N_d;大型模型此比較低(計算密集) |
吞吐指標
| 指標 | 說明 |
|---|---|
| TFLOPS/GPU | ZeRO 的額外通訊會降低有效吞吐,但節省的記憶體允許更大的 batch size,兩者可能互相抵消 |
| Model FLOPS Utilization (MFU) | 實際吞吐 / 硬體峰值理論吞吐;ZeRO-3 的 MFU 受通訊影響 |
供需與市場資料
需求側
| 驅動因素 | 資料/趨勢 |
|---|---|
| 模型規模增長 | 引數量從十億級→千億→萬億級趨勢未止,單卡裝不下已成常態 |
| 訓練成本壓力 | 大型模型訓練成本動輒數百萬至數千萬美元,記憶體最佳化直接影響所需 GPU 數量和訓練時長 |
| 微調需求爆發 | RLHF/SFT 等後訓練階段同樣需要全引數更新或大幅梯度計算 |
供給側
| 提供方 | 方案 | 定位 |
|---|---|---|
| Microsoft / DeepSpeed | ZeRO Stage 1/2/3 + Offload/Infinity | 最完整的 ZeRO 實現,開源 |
| Meta / PyTorch | FSDP(Fully Sharded Data Parallel) | ZeRO-3 等效,PyTorch 原生整合 |
| NVIDIA / Megatron-LM | 自有張量/流水線並行 | 與 ZeRO 互補而非替代 |
| Google / JAX | FSDP(JAX 版)/ GSPMD | JAX 生態下的自動分片 |
市場規模參考
大型模型訓練基礎設施市場規模尚無權威公開資料。根據公開資訊粗略估算:2024 年全球 AI 訓練 GPU 市場規模超過 300 億美元 [估算,來源:綜合 NVIDIA 財報及行業分析],ZeRO/FSDP 等記憶體最佳化技術是該市場的效率乘數——它們不直接產生營收,但決定了同樣的硬體預算能訓練多大的模型。
代表公司與資本對映
核心貢獻方
| 公司/團隊 | 產品/方案 | 關聯資產 | 說明 |
|---|---|---|---|
| Microsoft | DeepSpeed(含 ZeRO) | MSFT | ZeRO 發明者,DeepSpeed 是其 AI 基礎設施的關鍵元件,服務於 Azure AI 和 OpenAI 的訓練需求 |
| Meta | PyTorch FSDP | META | ZeRO-3 等效實現,PyTorch 生態覆蓋最廣 |
| NVIDIA | Megatron-LM(與 DeepSpeed 融合) | NVDA | 硬體+軟體棧,Megatron-DeepSpeed 組合廣泛使用 |
間接受益方
| 領域 | 說明 |
|---|---|
| GPU 網際網路絡 | ZeRO Stage 3 依賴高頻寬互連(NVLink, InfiniBand),利好網路裝置廠商 |
| CPU/記憶體/儲存 | ZeRO-Offload/Infinity 使 CPU 記憶體和 NVMe 參與訓練,對伺服器配置提出更高要求 |
| 大型模型公司 | OpenAI、Anthropic、Google DeepMind 等大型模型訓練方是 ZeRO 類技術的直接受益者 |
投資邏輯
核心觀點
-
ZeRO 是基礎設施層的”隱性標準”:不是獨立產品,而是嵌入 DeepSpeed/PyTorch FSDP 等架構的核心技術。投資機會在使用 ZeRO 的公司(訓練大型模型),而非 ZeRO 本身。
-
記憶體最佳化技術持續迭代:隨著模型規模繼續增長,ZeRO 的後續版本(++、Infinity 等)和其他新技術(如梯度檢查點、選擇性啟用重計算、CPU offload)將持續演進。這是一個持續投入、持續回報的技術方向。
-
與硬體協同演進:HBM 容量增長(HBM3e 36GB/stack → 未來更大)+ ZeRO 記憶體最佳化 = 更大的可訓練模型。兩者是乘法關係,而非替代關係。
-
競爭格局:DeepSpeed(Microsoft)和 PyTorch FSDP(Meta)是兩強格局。NVIDIA 在 Megatron 側有自己的並行策略但與 ZeRO 互補。Google 的 JAX 走 GSPMD 自動分片路線。
風險提示
- 如果未來硬體視訊記憶體足夠大(如 HBM 容量 × 數量遠超模型需求),ZeRO 的價值邊際遞減
- 新的訓練範式(如狀態空間模型、稀疏訓練)可能改變記憶體瓶頸的位置
- DeepSpeed 的 Microsoft 屬性可能影響其他公司的採納意願(FSDP 的中立性優勢)
常見誤讀糾偏
❌ 誤讀 1:ZeRO Stage 3 與張量並行(Tensor Parallelism)是同一回事
糾偏:兩者本質不同。
| 維度 | ZeRO Stage 3 | Tensor Parallel |
|---|---|---|
| 切分物件 | 資料並行維度的狀態(引數/梯度/最佳化器的冗餘副本) | 模型單層內的權重矩陣 |
| 通訊型別 | AllGather(拼回完整引數)+ ReduceScatter(聚合梯度) | AllReduce/AllGather(層內啟用/結果同步) |
| 適用範圍 | 任何模型(黑盒級別) | 需要修改模型程式碼,適配特定運算元(如 MLP、Attention) |
| 通訊頻率 | 每個 micro-batch 的前向/反向各一次 | 每層的前向/反向各一次 |
關鍵區別:ZeRO 切的是”冗餘”——每個 DP 程序的完整引數副本;Tensor Parallel 切的是”必須”——單層內權重矩陣的計算維度。
❌ 誤讀 2:ZeRO Stage 3 的通訊量是 Stage 1/2 的 1.5 倍,所以效能差很多
糾偏:
- 通訊量 ≠ 通訊時間。實際通訊時間取決於通訊/計算重疊程度。在大型模型中,計算時間遠大於通訊時間(計算密集),額外的 AllGather 可以被計算掩蓋(overlap)。
- 記憶體節省 → 更大 batch size → 更高 MFU。Stage 3 節省的記憶體可以用於增大 batch size,提升計算效率,部分甚至完全抵消額外通訊開銷。
- 通訊量是相對於模型引數 Φ 而言的,而非相對於總資料量。對於大型模型(Φ 很大),計算量(約 6Φ × sequence_length × batch_size)遠大於通訊量(3Φ),通訊/計算比很低。
❌ 誤讀 3:ZeRO 只適用於預訓練
糾偏:ZeRO 適用於任何需要梯度更新的分散式訓練場景:
- 全引數微調(Full Fine-tuning):與預訓練相同,ZeRO 直接適用
- LoRA / QLoRA:這些方法本身減少了可訓練引數,但 ZeRO 仍然可以最佳化基礎模型引數的記憶體佔用(即使 frozen 引數不更新,仍需視訊記憶體儲存)
- RLHF:PPO 訓練中同時存在 policy model、reference model、reward model、value model,記憶體壓力極大,ZeRO 可以最佳化每個模型的狀態記憶體
- 推論(有限適用):ZeRO-Infinity 的 offload 思路可用於推論時的模型載入
❌ 誤讀 4:FSDP 和 ZeRO-3 完全等價
糾偏:概念上等效(都是分片引數/梯度/最佳化器狀態),但實現和功能有差異:
- FSDP 是 PyTorch 原生 API,與 PyTorch 生態(編譯、自動混合精度等)深度整合
- ZeRO-3 (DeepSpeed) 有更多配置選項(offload、infinity、++等擴充套件),靈活性更高
- 兩者在具體通訊排程、記憶體碎片管理、activation checkpointing 策略等方面有工程差異
- 選擇取決於團隊技術棧偏好和具體需求
學習路徑
入門(1-2 小時)
- 閱讀本文
- 觀看 Microsoft DeepSpeed 官方 YouTube 講解(搜尋 “DeepSpeed ZeRO”)
- 理解混合精度訓練的記憶體構成
進階(1-2 天)
- 精讀原論文:Samyam Rajbhandari et al., “ZeRO: Memory Optimizations Toward Training Trillion Parameter Models”, OSDI 2020
- 動手實驗:用 DeepSpeed 在 2-4 張 GPU 上跑 ZeRO Stage 1/2/3 的示例,觀察視訊記憶體佔用變化
- 對比閱讀:PyTorch FSDP 官方教程
專家(持續)
- DeepSpeed 原始碼:
deepspeed/runtime/zero/目錄下的實現 - ZeRO-Offload/Infinity 論文:理解層次化儲存的工程挑戰
- Megatron-DeepSpeed 融合方案與萬億引數訓練實踐