🏠 总目录📚 本教程 ZeRO 与 FSDP
📑 本页目录(点开跳转)

11 · ZeRO 与 FSDP

28 分钟 | ⭐⭐ 最重要的一个分布式技术


🎯 一句话

数据并行的浪费很明显:N 张卡存着 N 份一模一样的优化器状态。 ZeRO 的想法简单到几乎是显然的:别存 N 份,切开,每卡存 1/N,用的时候再凑。 它让"数据并行"从省不了显存变成能训万亿参数。


🗑️ 一、被浪费的显存

   8 卡数据并行训练 7B 模型(第 3 章的账):

   每张卡:
   参数 14 GB + 梯度 14 GB + 优化器状态 84 GB = 112 GB
   × 8 卡 = 896 GB

   ⭐ 但真正需要的信息只有 112 GB —— 剩下 784 GB 全是【完全一样的副本】

🔑 ZeRO 的核心洞察这些副本在大部分时间里都是闲置的。 优化器状态只在 step() 那一瞬间用到 —— 既然如此,为什么每张卡都要一直存着完整的一份?


📊 二、三个 Stage

Stage 0(普通 DDP)
  • 参数[完整] 梯度[完整] 优化器[完整] 112 GB/卡
Stage 1:切【优化器状态】
  • 参数[完整] 梯度[完整] 优化器[1/N] ~38.5 GB/卡
Stage 2:+ 切【梯度】
  • 参数[完整] 梯度[1/N] 优化器[1/N] ~26.3 GB/卡
Stage 3:+ 切【参数】⭐
  • 参数[1/N] 梯度[1/N] 优化器[1/N] ~14 GB/卡

(8 卡,7B 模型)

Stage 显存倍数 通信量 vs DDP 什么时候用
1 ÷4 (一样)⭐ 几乎白送,默认就该开
2 ÷8 (一样)⭐ 性价比最高 ⭐⭐
3 ÷N 1.5× 模型实在装不下时

🔑 最反直觉也最重要的一点Stage 1 和 2 的通信量和普通 DDP 完全一样。

为什么?回忆第 5 章那个恒等式: AllReduce = ReduceScatter + AllGather

   DDP:     AllReduce(梯度)                    = RS + AG
   ZeRO-2:  ReduceScatter(梯度) → 各自更新自己那份
             → AllGather(参数)
             = RS + AG   ⭐ 一模一样!

ZeRO-2 只是把一次 AllReduce 拆成两半,在中间插入了"各自更新自己那片"。 所以它是几乎免费的显存节省。


🔄 三、Stage 3 怎么工作(FSDP 的核心)

   平时:每张卡只存 1/N 的参数

   前向传播到第 L 层时:
   ① AllGather:临时凑出【第 L 层】的完整参数  ⭐
   ② 算这一层
   ③ 立刻【丢掉】刚凑出来的那份            ⭐
   ④ 继续下一层

   反向同理,再加一次 ReduceScatter 汇总梯度
   ⭐ 关键:任何时刻只有【一层】的完整参数在显存里
     → 峰值显存 ≈ 分片参数 + 最大单层的完整参数

💡 这就是为什么它叫 Fully Sharded Data Parallel: 从"每卡一份完整模型"变成"每卡一片,用时凑齐"。 代价是多了 1.5 倍的通信 —— 但换来的是 N 倍的显存。

⭐ 预取:让通信藏起来

   朴素做法:算第 L 层时才去 AllGather 第 L 层 → GPU 干等 💀

   ✅ 预取(prefetch):
     算第 L 层的【同时】,后台 AllGather 第 L+1 层的参数
     → 通信和计算重叠

   ⭐ 这是 FSDP 性能好坏的关键。配置没调好,可能比 DDP 还慢。
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy, MixedPrecision
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
import functools

model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,   # ⭐ = ZeRO-2
    # ShardingStrategy.FULL_SHARD                       #   = ZeRO-3
    # ShardingStrategy.HYBRID_SHARD                     #   机内切、跨机复制 ⭐
    auto_wrap_policy=functools.partial(
        transformer_auto_wrap_policy,
        transformer_layer_cls={TransformerBlock},       # ⭐ 按层包,非常重要
    ),
    mixed_precision=MixedPrecision(
        param_dtype=torch.bfloat16,
        reduce_dtype=torch.bfloat16,      # ⭐ 通信也用 BF16,通信量减半
        buffer_dtype=torch.bfloat16,
    ),
    device_id=torch.cuda.current_device(),
)

⚠️ auto_wrap_policy 不设置是最常见的错误: 不设置的话整个模型被当成一个单元,AllGather 时要凑出全部参数 —— 显存节省几乎为零,而且会 OOM。 必须按 Transformer 层来包。

💡 HYBRID_SHARD:大规模的最优解

   问题:ZeRO-3 跨机 AllGather 参数,而跨机带宽只有 NVLink 的 1/15

   ✅ HYBRID_SHARD:
   · 【机内】8 卡做 FULL_SHARD(NVLink 快,切得彻底)
   · 【机间】做普通数据并行(只 AllReduce 梯度)

   → 显存降到 1/8,而跨机通信量和普通 DDP 一样 ⭐

⚖️ 四、ZeRO vs 张量并行:怎么选

ZeRO-3 / FSDP 张量并行
改代码 ✅ 几乎不用 ⚠️ 要改模型结构
通信频率 每层一次 AllGather 每层两次 AllReduce(更频繁)
能跨机吗 ✅ 可以(HYBRID 更好) 基本不行第 5 章
单层装不下时 ❌ 救不了 ✅ 能救 ⭐
适合 大多数场景 单层特别大 / 超大模型

实用建议的顺序

   ① 先上 ZeRO-2(几乎免费)
   ② 不够 → ZeRO-3 / FSDP FULL_SHARD
   ③ 多机 → HYBRID_SHARD
   ④ 单层都装不下 → 加张量并行(第 12 章)
   ⑤ 还不够 → 加流水线并行(第 13 章)

⚠️ 五、四个坑

症状 解法
没设 auto_wrap_policy 显存没省,还 OOM 按 Transformer 层包 ⭐
预取没配好 FSDP 比 DDP 还慢 backward_prefetch=BACKWARD_PRE
保存 checkpoint 时 OOM 保存的瞬间要凑齐全部参数 SHARDED_STATE_DICT 分片保存 ⭐
和梯度检查点叠加顺序错 报错或不生效 先 checkpoint_wrapper,再 FSDP
# ⭐ 分片保存:每卡只存自己那片,不需要凑齐
from torch.distributed.checkpoint.state_dict import get_state_dict
import torch.distributed.checkpoint as dcp

dcp.save({"model": model, "optim": opt}, checkpoint_id="ckpt/step1000")

💥 "保存 checkpoint 时 OOM" 是一个非常常见的翻车点: 训练跑得好好的,一到保存就炸 —— 因为默认的 FULL_STATE_DICT 会把所有分片汇总到 rank 0。大模型必须用分片保存。


🔗 和站内其他章的关系

相关的地方 这里的位置
第 3 章 优化器状态 84GB ZeRO 要切的就是它
第 5 章 AllReduce=RS+AG ZeRO-2 免费的原因 ⭐⭐
第 10 章 DDP 一块显存都没省 本章是它的修正
第 5 章 跨机带宽是瓶颈 HYBRID_SHARD 的动机
第 9 章 决策树 ZeRO 在其中的位置
《机器学习与深度学习基础》09 优化器与学习率 Adam 的一阶/二阶动量 Stage 1 切的就是这两份状态 —— 不知道 Adam 存了什么,就看不懂 ZeRO 省在哪 ⭐⭐
《强化学习基础》12 RLHF 全流程 PPO 同时装四个模型 显存是 SFT 的 4 倍以上,是 Stage 3 + CPU offload 最典型的用武之地 ⭐
《Kaggle竞赛方法论》05 预训练模型详解 LoRA / PEFT 另一条路:LoRA 让优化器状态少两个数量级;小规模微调先试它,别急着上 ZeRO

✅ 检查点

  1. ZeRO 的核心洞察是什么?8 卡训 7B 浪费了多少显存?
  2. 三个 Stage 各切什么?显存各降多少?
  3. 为什么 ZeRO-1 和 2 的通信量和 DDP 完全一样?(说出那个恒等式)
  4. Stage 3 前向传播时怎么工作?峰值显存是多少?
  5. 什么是预取?为什么它是 FSDP 性能的关键?
  6. auto_wrap_policy 不设置会怎样?
  7. HYBRID_SHARD 解决什么问题?
  8. ZeRO-3 和张量并行的关键区别?什么时候必须用后者?
  9. "保存 checkpoint 时 OOM"是怎么回事?
👀 答案
  1. 洞察:这些副本大部分时间都闲置(优化器状态只在 step() 那一瞬间用到),既然如此为什么每卡都要一直存完整的一份。8 卡训 7B 共占 896GB,真正需要的只有 112GB,784GB 是完全一样的副本
  2. Stage 1 切优化器状态(÷4)、Stage 2 再切梯度(÷8)、Stage 3 再切参数(÷N)。
  3. 因为 AllReduce = ReduceScatter + AllGather。DDP 做一次 AllReduce(梯度);ZeRO-2 做 ReduceScatter(梯度) → 各自更新自己那份 → AllGather(参数),总量一模一样——它只是把一次 AllReduce 拆成两半,中间插入"各自更新自己那片"。
  4. 算第 L 层时 AllGather 临时凑出该层完整参数 → 算完立刻丢掉 → 继续下一层。峰值 ≈ 分片参数 + 最大单层的完整参数(任何时刻只有一层的完整参数在显存里)。
  5. 算第 L 层的同时后台 AllGather 第 L+1 层,让通信和计算重叠。关键是因为不预取的话 GPU 会干等 AllGather,配置没调好可能比 DDP 还慢
  6. 整个模型被当成一个单元,AllGather 时要凑出全部参数 → 显存节省几乎为零而且会 OOM。必须按 Transformer 层包。
  7. 解决ZeRO-3 跨机 AllGather 参数而跨机带宽只有 NVLink 1/15 的问题。做法:机内 FULL_SHARD、机间普通数据并行 → 显存降到 1/8,跨机通信量和普通 DDP 一样
  8. ZeRO-3 几乎不用改代码、能跨机;张量并行要改模型结构、基本不能跨机,但单层装不下时只有它能救
  9. 默认的 FULL_STATE_DICT把所有分片汇总到 rank 0,保存瞬间要凑齐全部参数。解法:用 SHARDED_STATE_DICT / dcp.save 分片保存

🛑 可以停在这里

走神救援

核心洞察:N 张卡存着 N 份一模一样的优化器状态,而它们只在 step() 那一瞬间用到(8 卡训 7B 共 896GB,真正需要 112GB,784GB 是副本)。三个 Stage:Stage1 切优化器状态(÷4)、Stage2 再切梯度(÷8)、Stage3 再切参数(÷N)。⭐⭐最重要也最反直觉:Stage 1 和 2 的通信量和 DDP 完全一样——因为 AllReduce = ReduceScatter + AllGather,ZeRO-2 只是把一次 AllReduce 拆成两半、中间插入"各自更新自己那片"几乎免费的显存节省,默认就该开Stage 3(FSDP):平时每卡只存 1/N,算第 L 层时 AllGather 临时凑出该层完整参数 → 算完立刻丢掉,峰值 ≈ 分片参数 + 最大单层;代价 1.5× 通信。⭐预取是性能关键(算第 L 层时后台 AllGather 第 L+1 层,不预取 GPU 会干等,配置没调好可能比 DDP 还慢)。⚠️⭐auto_wrap_policy 不设置是最常见错误——整个模型当成一个单元,显存节省几乎为零还会 OOM,必须按 Transformer 层包。⭐HYBRID_SHARD 是多机最优解:机内 FULL_SHARD(NVLink 快)+ 机间普通数据并行 → 显存降 1/8 而跨机通信量和 DDP 一样vs 张量并行:ZeRO 几乎不用改代码、能跨机;张量并行要改结构、基本不能跨机,但单层装不下时只有它能救。⭐顺序:ZeRO-2 → ZeRO-3 → HYBRID → 加张量并行 → 加流水线并行。💥⚠️"保存 checkpoint 时 OOM"是常见翻车点——默认 FULL_STATE_DICT 把所有分片汇总到 rank 0,大模型必须用 dcp.save 分片保存。还有:reduce_dtype=bfloat16 让通信量减半、梯度检查点要先 checkpoint_wrapper 再 FSDP

下一节 👉 12-张量并行.md

打卡记录保存在你的浏览器里,首页能看到总进度