7.2 ZeRO 优化:如何突破单卡显存限制
ZeRO(Zero Redundancy Optimizer)由 Microsoft DeepSpeed 团队提出,是解决数据并行显存冗余问题的核心技术。7.1 末尾那 1.12 TB 之所以是死路,不在于它大,而在于每张数据并行卡上都躺着一份一模一样的;而 7.1.3 留下的那个恒等式——AllReduce 等于 Reduce-Scatter 加 All-Gather——已经把拆开它的入口露在了外面。
7.2.1 冗余分析
在标准数据并行中,每张 GPU 都存储了完整的模型状态——参数、梯度和优化器状态。冗余的关键不在于“每张卡只需要自己数据的梯度”,因为每个 rank 逻辑上仍要参与完整模型的梯度同步;真正的浪费是这些模型状态在每个数据并行副本上都被完整复制了一份。ZeRO 的目标就是把优化器状态、梯度和参数按数据并行 rank 分片保存,同时保持等价的数据并行更新语义。
以一个 15 亿参数的模型(约 GPT-2 规模)为例,使用 FP16 训练和 Adam 优化器:
FP16 参数
3 GB
字节
FP16 梯度
3 GB
与参数同大小
FP32 优化器状态
18 GB
参数副本 + 一阶矩 + 二阶矩
总计
24 GB
每张 GPU
表 7-1:每张 GPU 的显存占用明细
在 64 张 GPU 上做数据并行时,这 24 GB 在每张卡上都有一份完整拷贝——总计 1536 GB 的显存中只有 24 GB 是唯一状态,剩余 1512 GB(63/64)都是冗余拷贝。
下表固定参数状态 3 GB、梯度 3 GB、Adam 优化器状态 18 GB,并按 逐项分片。它只核算持久模型状态,不包含激活、临时 AllGather buffer、通信 workspace 和碎片。
DDP
3
3
18
64
24.000000
1.00
ZeRO-1
3
3
18
64
6.281250
3.82
ZeRO-2
3
3
18
64
3.328125
7.21
ZeRO-3
3
3
18
64
0.375000
64.00
表 7-2:64 路数据并行下的模型状态分片重算。ZeRO-1 为 ,ZeRO-2 为 ,ZeRO-3 为 ;“相对 DDP”是 ,不是端到端吞吐倍数。
7.2.2 ZeRO 的三个阶段
ZeRO 通过将模型状态分片(Shard)到多张 GPU 上来消除冗余:
ZeRO-1(优化器状态分片):每张 GPU 只持有 的优化器状态。更新参数时,各卡负责更新自己那部分参数,然后通过 AllGather 同步更新后的参数。显存减少约 4 倍。
ZeRO-2(梯度分片):在 ZeRO-1 基础上,梯度也分片存储。每张 GPU 只保留与自己负责的参数对应的梯度,其余在 Reduce-Scatter 后丢弃。显存进一步减少约 2 倍。
ZeRO-3(参数分片):最彻底的方案——连参数本身也分片存储。每张 GPU 只持有 的参数,在前向和反向传播时通过 AllGather 临时获取完整参数。显存减少与 GPU 数量成正比。
7.2.3 通信与效率权衡
ZeRO 通过改变通信模式来换取显存节省:
标准 DDP
1×
AllReduce
1×
基准
ZeRO-1
~4×
AllGather
1×
几乎无
ZeRO-2
~8×
Reduce-Scatter + AllGather
1×
几乎无
ZeRO-3
线性于 GPU 数
AllGather(前/后向)
~1.5×
有一定开销
ZeRO-1/2 虽然使用了不同于 DDP 的通信原语(AllGather 和 Reduce-Scatter 而非 AllReduce),但通过精巧的通信调度,总的通信数据量与 DDP 相同,因此几乎没有性能损失。这是通过计算-通信重叠和梯度分片感知的通信顺序实现的。
ZeRO-3 需要额外的 AllGather 操作来在前向和反向传播中获取完整参数,通信量增加约 50%。但 DeepSpeed 通过参数预取和流水线化重计算等优化将实际开销控制在可接受的范围内(通常性能下降 10-20%)。
ZeRO 的出现使得数据并行能够训练远超单卡显存容量的模型,成为了大模型训练的基础设施级技术。PyTorch 原生的 FSDP(Fully Sharded Data Parallel)正是 ZeRO-3 全分片思想的官方实现,7.7 节谈分片检查点与拓扑可迁移检查点时还会再遇到它。
FSDP2:从「拍平成一维」到「逐参数分片」
FSDP 本身也演进过一轮,而这轮演进的动机不在显存,在可组合性。第一代 FSDP 把一组参数拍平(flatten)成一个一维缓冲区再切分——这对通信是高效的,但「原来的参数」在实现里已经不存在了,任何想要按参数区别对待的需求都会变得别扭:给某几个参数单独设混合精度、只量化一部分权重、冻结其中一些、或者导出一份按参数组织的状态字典。
PyTorch 的第二代实现 fully_shard(FSDP2)改用 DTensor 逐参数分片:官方文档的表述是「针对高性能 eager 模式、同时用逐参数分片改善易用性」。它的用户契约相当直白——初始化时把 model.parameters() 原地从普通 torch.Tensor 转成 DTensor;前向/反向之前由 pre-hook 负责 all-gather 参数并转回普通张量;前向/反向之后由 post-hook 释放未分片的参数(这一步不需要通信)。参数始终保留自己的身份,切分信息记录在 DTensor 的 placement 里,于是上面那些「按参数区别对待」的需求重新变得自然。
同一套 device mesh 抽象还顺带把 HSDP(Hybrid Sharded Data Parallel)表达了出来:mesh 是 1D 时参数在这一维上全分片,placement 为 (Shard(0),);mesh 是 2D 时参数在第 1 维分片、在第 0 维复制,placement 为 (Replicate(), Shard(0))。这正是「节点内分片、节点间复制」的形式化——它的适用场景来自 7.4.3 的同一条原则:全分片的通信量随分片组变大而变大,当分片组已经跨过慢速互连时,把复制维放到慢链路、把分片维留在快链路,总代价反而更低。
从 ZeRO 论文到 FSDP2,分片的数学一步没变,变的是分片这件事怎么表达——而正是这个表达方式决定了它能不能和量化、混合精度、张量并行干净地叠在一起。
最后更新于
