ZeRO 和 FSDP 是什么?它们如何降低训练显存?
简化版
普通数据并行里,每张卡都存一份完整的模型参数、梯度和优化器状态,极其浪费。ZeRO(Zero Redundancy Optimizer) 的思想是把这些状态”分片”到各张卡上,谁用到谁临时聚合,消除冗余,从而在同样的卡上训练大得多的模型。它分三级:ZeRO-1 分片优化器状态、ZeRO-2 再分片梯度、ZeRO-3 连参数也分片。FSDP(Fully Sharded Data Parallel) 是 PyTorch 官方实现的类似方案,等价于 ZeRO-3——把模型完全分片,前向/反向时按需 all-gather 聚合、用完即释放。
详细版
问题:数据并行的显存冗余
用 Adam + 混合精度训练时,每个参数要占的「模型状态」显存约为(常见估算):
FP16 参数 2B + FP16 梯度 2B + FP32 参数副本 4B + FP32 动量 m 4B + FP32 方差 v 4B ≈ 16 字节/参数
所以一个 7.5B 模型光「模型状态」就约 120GB。在普通数据并行下,N 张卡各存一份这 120GB,冗余 N 倍,卡再多也放不下更大的模型。
ZeRO 三级分片
| 级别 | 分片内容 | 显存节省(DP 度 = N) | 通信开销 |
|---|---|---|---|
| ZeRO-1 | 优化器状态(m、v、FP32 副本) | 约 4× | 与 DP 相当 |
| ZeRO-2 | + 梯度 | 约 8× | 与 DP 相当 |
| ZeRO-3 | + 参数 | 约 N×(线性) | 略增(需 all-gather 参数) |
- ZeRO-1:优化器状态(占大头,约 12B/参数)分片到各卡,每卡只存 1/N。
- ZeRO-2:梯度也分片,进一步省。
- ZeRO-3:连参数本身也分片,每卡平时只存 1/N 的参数,用到某层时临时 all-gather 出完整参数、算完立即释放。显存随卡数近似线性下降,能训最大的模型。
FSDP
- PyTorch 原生的全分片数据并行,机制等价 ZeRO-3:参数、梯度、优化器状态全分片。
- 前向到某层时 all-gather 聚合该层完整参数 → 计算 → 释放;反向同理,再 reduce-scatter 分发梯度。
- 是当前 PyTorch 生态训练大模型的主流选择。
完整版教学
一、先看清冗余到底浪费在哪
标准数据并行的最大问题是重复存储。N 张卡组成数据并行组,每张卡都持有:完整参数、完整梯度、完整优化器状态。这三样里,优化器状态最占地方——用 Adam 时,每个参数要额外存 FP32 的一阶动量 m、二阶动量 v,加上 FP32 参数副本,混合精度下「模型状态」高达约 16 字节/参数,其中优化器相关就占约 12 字节。
关键洞察:数据并行里,每张卡其实只负责更新”一部分”参数就够了,没必要每张卡都存全套。 ZeRO 就是把「本可以分工的东西」真正分给各卡,消除冗余。
二、ZeRO 的核心思想:分片 + 按需聚合
ZeRO(Zero Redundancy Optimizer)不改变数据并行「各算各的数据、同步梯度」的本质,而是把模型状态沿卡分片存储,需要完整数据时临时用集合通信(all-gather / reduce-scatter)拼出来,用完就丢。
- 平时:每张卡只存自己负责的那 1/N 份状态,显存大降。
- 用时:需要完整参数/梯度时,临时从各卡聚合,算完立即释放,不长期占显存。
这是典型的「用一点通信,换大量显存」——而带宽往往比显存更容易扩展,所以非常划算。
三、三级分片:省得越多,通信越多
ZeRO 分三级,逐级把更多东西分片:
ZeRO-1(分片优化器状态):把 Adam 的 m、v、FP32 副本分片到各卡。因为优化器状态占大头(约 12B/参数),这一级就能省约 4 倍显存,而通信几乎不增加(只是更新时各卡管自己那份,再同步)。性价比最高、最常用。
ZeRO-2(再分片梯度):梯度也分片,每卡只保留它负责参数的梯度。省到约 8 倍,通信仍与普通 DP 相当(把 all-reduce 换成等价的 reduce-scatter + all-gather)。
ZeRO-3(连参数也分片):最激进,参数本身也分片,每卡平时只存 1/N 的参数。前向/反向到某层时,all-gather 临时拼出该层完整参数 → 计算 → 立即释放。显存随卡数近似线性下降,理论上卡越多能训越大的模型。代价是多了参数聚合的通信(每层前向反向都要 all-gather 参数)。
记忆阶梯:ZeRO-1 分优化器状态 → ZeRO-2 加分梯度 → ZeRO-3 再分参数。省显存能力递增(4×→8×→N×),通信开销也递增。
四、FSDP:PyTorch 版的 ZeRO-3
FSDP(Fully Sharded Data Parallel)是 PyTorch 官方的全分片实现,机制等价于 ZeRO-3:
- 模型按「分片单元」(如每个 Transformer 层)拆开,参数、梯度、优化器状态全部分片;
- 前向:进入某层时
all-gather聚合出该层完整参数 → 计算 → 用完立即释放显存; - 反向:同样 all-gather 参数算梯度,再用
reduce-scatter把梯度分发回各自负责的卡; - 各卡用本地那份梯度更新本地那份参数。
FSDP 已成为 PyTorch 训练大模型的主流工具,配合混合精度、梯度检查点、激活分片等,能在有限卡数上训很大的模型。DeepSpeed 的 ZeRO 是另一套成熟实现,二者思想一致。
五、ZeRO 和三种并行是什么关系
初学者容易混淆 ZeRO 和张量/流水线并行。关键区分:
- ZeRO / FSDP 本质是”去冗余的数据并行”:它没有把单层计算拆到多卡(那是张量并行),也没按层分段(那是流水线并行)。每张卡仍然算完整的前向/反向,只是参数/梯度/优化器状态的存储被分片了。
- 因此 ZeRO 可以和 TP、PP 叠加:3D 并行解决”计算怎么拆”,ZeRO 解决”数据并行维度上的存储冗余”。
简要说:TP/PP 拆的是”计算”,ZeRO 拆的是”存储”,两者正交互补。
六、面试拆解算例
预训练题最好落到预算账和稳定性账。假设训练一个 7B 模型,目标 token 数是 1T,按常见粗估 6 × 参数量 × token 数,训练计算量约为 6 × 7e9 × 1e12 = 4.2e22 FLOPs。如果有效集群算力是 1e18 FLOPs/s,理想情况下也要约 42,000 秒;现实还要扣通信、数据加载、checkpoint 和故障恢复的损耗。
training_flops ≈ 6 * N * D
N = 7e9 parameters
D = 1e12 tokens
training_flops ≈ 4.2e22
| 账本 | 关键变量 | 常见瓶颈 | 排查信号 |
|---|---|---|---|
| 数据账 | token 数、重复率、质量分 | 脏数据和污染 | eval 异常偏高 |
| 算力账 | GPU 数、利用率、通信 | MFU 低 | step time 抖动 |
| 显存账 | batch、序列长、优化器状态 | OOM | 激活占用过高 |
| 稳定性账 | 学习率、精度、梯度 | loss spike | overflow/NaN |
语料 -> 清洗去重 -> tokenization -> 分布式训练 -> checkpoint -> 评测
| | | | |
质量 覆盖率 吞吐 可恢复 能力验证
所以回答「ZeRO 和 FSDP 是什么?它们如何降低训练显存?」时,不能只说某个技巧“省显存”或“加速”。要说明它省的是哪一笔账、牺牲了什么、线上训练日志里应该观察哪个信号。
七、常见误区与追问
- 误区:ZeRO 是一种模型并行(拆计算)。 不是,它是分片存储的数据并行,计算仍在每卡完整进行。
- 误区:优化器状态不占多少显存。 恰恰相反,Adam 的 m/v + FP32 副本约占 12B/参数,是显存大头,所以 ZeRO-1 先拿它开刀。
- 追问:为什么每参数约 16 字节? 混合精度 Adam:FP16 参数 2 + FP16 梯度 2 + FP32 参数副本 4 + FP32 动量 4 + FP32 方差 4 = 16。
- 追问:ZeRO-3 为什么能近线性省显存? 参数/梯度/优化器状态全按卡数分片,每卡只存 1/N,用时临时聚合即释放。
- 追问:FSDP 和 ZeRO-3 区别? 思想等价,FSDP 是 PyTorch 原生实现,ZeRO 是 DeepSpeed 实现。
- 追问:ZeRO-3 的代价? 每层前向/反向都要 all-gather 参数,通信量增加,需高带宽。
- 追问:ZeRO 和张量并行冲突吗? 不冲突,正交——一个分存储、一个分计算,常一起用。
八、加强记忆
ZeRO 记「去冗余、三级分片」:普通数据并行每卡都存整份参数/梯度/优化器状态(约 16B/参数,优化器占大头),ZeRO 把它们分片、用时聚合、用完释放——ZeRO-1 分优化器状态(省~4×)→ ZeRO-2 加分梯度(~8×)→ ZeRO-3 再分参数(~N× 近线性)。FSDP = PyTorch 版 ZeRO-3,前向按层 all-gather 参数、算完释放、reduce-scatter 回传梯度。最关键的辨析:ZeRO 拆的是”存储”(仍是数据并行),TP/PP 拆的是”计算”,两者正交可叠加。用「一点通信换大量显存」这句话锚住它的本质。