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