🏠 总目录📚 本教程 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 1×(一样)⭐ 几乎白送,默认就该开
2 ÷8 1×(一样)⭐ 性价比最高 ⭐
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. 平时:每张卡只存 1/N 的参数
  2. 前向传播到第 L 层时:
  3. AllGather:临时凑出【第 L 层】的完整参数 ⭐
  4. 算这一层
  5. 立刻【丢掉】刚凑出来的那份 ⭐
  6. 继续下一层
  7. 反向同理,再加一次 ReduceScatter 汇总梯度

算一算

⭐ 关键:任何时刻只有【一层】的完整参数在显存里

→ 峰值显存 ≈ 分片参数 + 最大单层的完整参数

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

⭐ 预取:让通信藏起来

结果对照

朴素做法:算第 L 层时才去 AllGather 第 L 层→GPU 干等 💀
✅ 预取(prefetch):
算第 L 层的【同时】,后台 AllGather 第 L+1 层的参数
通信和计算重叠
⭐ 这是 FSDP 性能好坏的关键。配置没调好,可能比 DDP 还慢。
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
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
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# ⭐ 分片保存:每卡只存自己那片,不需要凑齐
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 份一模一样的优化器状态,而它们只在参数更新那一瞬间用得到。 绝大部分显存是副本。

三个 Stage 逐级切:先切优化器状态、再切梯度、最后切参数。

⭐⭐ 最重要也最反直觉的一条:前两级的通信量和普通数据并行完全一样。 因为 AllReduce 本来就等于 ReduceScatter 加 AllGather——它只是把一次 AllReduce 拆成两半、中间插进「各自更新自己那一片」。⭐ 于是显存节省几乎是免费的,默认就该开。

第三级平时每卡只存一片,用到某一层时临时凑出完整参数、算完立刻丢——⭐ 峰值降到「分片参数加最大的那一层」,代价是通信量涨一半。

⭐ 预取是性能关键:算这一层的时候后台去取下一层,⚠️ 不预取 GPU 就干等,配置没调好可能比不切还慢。 ⚠️⭐ 最常见的错误是没设包装策略——整个模型被当成一个单元,显存几乎没省下来还会 OOM;必须按 Transformer 层包。

⭐ 多机的最优解是混合切分:机内全切(走高速互联)、机间用普通数据并行——显存降下来了,而跨机通信量和不切时一样。

⭐ 和张量并行的分工:这条路几乎不用改代码、能跨机;张量并行要改结构、基本不能跨机,⭐ 但单层都装不下时只有它能救。

⭐ 上手顺序:先开第二级 → 第三级 → 混合 → 再叠张量并行 → 最后流水线并行。 💥⚠️ 一个常见翻车点:保存 checkpoint 时 OOM——默认会把所有分片汇总到一张卡上,大模型必须用分片保存。

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

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