AIGC 基本功|数据并行与 ZeRO 显存切分-ZeRO

人工智能炼丹君
2026-09-05 / 0 评论 / 7 阅读 / 正在检测是否收录...

数据并行与 ZeRO 显存切分

所属方向:分布式训练 | 难度:进阶 | 前置知识:无
关键词:数据并行、DDP、ZeRO、优化器状态切分、FSDP、显存占用估算


01. 为什么需要它

先算一笔会让很多人意外的账:训练一个 75 亿参数模型,模型权重明明只有 15 GB,为什么放进 80 GB 显存的 GPU 还会爆?

假设用混合精度 Adam 训练。每个参数除了 2 字节的 FP16/BF16 权重,还要留下 2 字节梯度、4 字节 FP32 主权重、4 字节一阶动量和 4 字节二阶动量。合计不是 2 字节,而是 16 字节/参数:

$$7.5\times10^9\times16\ \mathrm{bytes}=120\ \mathrm{GB}$$

这 120 GB 还没有算激活、临时通信缓冲区、CUDA context 和显存碎片。换句话说,模型甚至没开始处理一帧视频,单是“为了训练而保存的状态”就已经放不下。

自然的反应是:“那我上 64 张卡。”可如果只用普通数据并行,每张卡都会保留完整的参数、完整的梯度和完整的 Adam 状态。64 张卡只是把不同 mini-batch 同时算得更快,每张卡看到的仍然是同一份 120 GB。总显存从 80 GB 变成 5120 GB,但单卡瓶颈一点没松,模型照样启动不了。

这就是 ZeRO 要解决的矛盾:数据并行已经拥有整个集群的总显存,却因为每张卡都复制同样的训练状态,只能使用其中一张卡的容量。ZeRO 不改变模型的数学计算,而是按顺序切掉三种冗余副本:

  • ZeRO-1 切优化器状态;
  • ZeRO-2 再切梯度;
  • ZeRO-3 连参数也切,只在某一层即将计算时临时拼回来。

把刚才的 75 亿参数模型放到 64 个数据并行 rank 上,三阶段的模型状态显存分别是 31.41 GB、16.64 GB 和 1.88 GB。最后一个数字正好是 $120/64$。这些数字不是宣传页上的模糊“最高节省多少”,而是本文后面会逐项推出来、再用代码复现的结果。

对视频生成尤其要补一句:视频 DiT 的 token 数很长,激活往往比模型状态还大。ZeRO 能拆掉模型状态冗余,但不会自动消灭激活。先分清是哪一类显存爆了,再决定用 ZeRO、激活重计算还是序列并行,比盲目把配置改成 stage 3 更重要。

02. 最小可用理解

三句话先建立框架:

  1. 数据并行让每个 rank 保存同一个模型、读取不同数据,反向后把局部梯度求平均,因而数学上等价于在合并后的大 batch 上做一次同步更新。
  2. 既然所有 rank 最终拿到相同梯度,它们各自保存一整套 Adam 状态并重复做同一份更新就是冗余;ZeRO 把状态分片,每个 rank 只负责其中 $1/N$。
  3. 切得越彻底,常驻显存越少,但参数就越需要“用前聚合、用后释放”,因此省下的显存会转化成通信、调度和峰值控制问题。

如果只记一件事:DDP 是复制计算,ZeRO 是切分状态;ZeRO 没有改变梯度本身。

这里的 rank 可以先理解为一个独立训练进程,通常一个 rank 独占一张 GPU。$N$ 表示数据并行进程数,world_size 就是 $N$。每个 rank 处理本地 batch,但每一步结束后所有 rank 必须得到一致参数,否则下一步就不再是同一个模型。

03. 数学推导

3.1 为什么平均梯度等价于一个大 batch

设全局 batch 有 $B$ 个样本,均匀切给 $N$ 个 rank,每个 rank 得到 $b=B/N$ 个样本。模型参数记作 $w$,第 $i$ 个样本的损失是 $\ell(x_i;w)$。

第 $r$ 个 rank 的局部平均损失为:

$$L_r(w)=\frac{1}{b}\sum_{i\in\mathcal{B}_r}\ell(x_i;w)$$

其中 $\mathcal{B}_r$ 是 rank $r$ 拿到的样本集合,$b$ 是本地 batch size。对参数求导,得到局部梯度:

$$g_r=\nabla_w L_r(w)=\frac{1}{b}\sum_{i\in\mathcal{B}_r}\nabla_w\ell(x_i;w)$$

同步数据并行对所有局部梯度取平均:

$$g=\frac{1}{N}\sum_{r=1}^{N}g_r=\frac{1}{Nb}\sum_{r=1}^{N}\sum_{i\in\mathcal{B}_r}\nabla_w\ell(x_i;w)=\frac{1}{B}\sum_{i=1}^{B}\nabla_w\ell(x_i;w)$$

右边正是把所有样本放在一张卡上算全局平均损失所得到的梯度。等价成立依赖三个前提:损失能按样本分解且样本计算不因分组改变(如本地 BatchNorm 统计会破坏这一条件);各 rank 从同一个参数 $w$ 出发;样本权重一致;更新前完成同步。若最后一个 batch 在各 rank 上不等长,仍然机械地“先本地平均再按 rank 平均”,样本就会被赋予错误权重。这也是分布式数据加载器必须谨慎处理尾 batch 的原因。

DDP 的核心任务由此非常明确:它不替你切输入,而是在 autograd 产生梯度后,把各 rank 的 $g_r$ 聚合成相同的 $g$。

3.2 先把 16 字节拆明白

设模型有 $P$ 个可训练参数,采用低精度前后向和 FP32 Adam 更新。每个参数对应:

  • 低精度参数:2 字节;
  • 低精度梯度:2 字节;
  • FP32 主参数:4 字节;
  • Adam 一阶矩 $m$:4 字节;
  • Adam 二阶矩 $v$:4 字节。

后面三项统称优化器状态,共 12 字节。普通 DDP 每个 rank 的模型状态显存是:

$$M_{\mathrm{DDP}}=(2+2+12)P=16P$$

注意这是一个状态下界模型,没有包含激活等剩余显存。它的价值不是预测 nvidia-smi 到个位数,而是先回答“哪一类状态值得切”。

实际套公式时,$P$ 应取需要优化器更新的参数量,不是配置文件里笼统写出的模型总参数量。冻结的视觉编码器通常仍要保存推理权重,却不产生梯度和 Adam 状态;共享 embedding 在参数统计中也只能算一次。优化器同样会改常数:SGD 无动量时几乎没有 $m,v$,8-bit Adam 会压缩这两份状态,某些纯 BF16 配置也不保留 FP32 主参数。所以正确动作不是背住“训练恒等于 16 字节/参数”,而是先列出当前训练栈真正持有的张量,再把每一项的元素数乘 dtype 字节数。本文使用 16 字节,是为了与 ZeRO 原论文的混合精度 Adam 假设严格对齐。

3.3 ZeRO 三阶段到底切了什么

ZeRO-1:只切优化器状态。 参数和梯度仍然每卡各有一份,12 字节的 FP32 主参数与 Adam 两个矩按 $N$ 份切开:

$$M_{\mathrm{Z1}}=4P+\frac{12P}{N}$$

当 $N$ 很大,第二项趋近于 0,但第一项中的 2 字节参数和 2 字节梯度仍然复制,所以 ZeRO-1 的极限是 $4P$,相对 $16P$ 最多节省 4 倍。

ZeRO-2:再切梯度。 常驻的完整副本只剩 2 字节低精度参数;梯度与优化器状态合计 14 字节一起分片:

$$M_{\mathrm{Z2}}=2P+\frac{14P}{N}$$

当 $N$ 很大,极限是 $2P$,相对普通 DDP 最多节省 8 倍。这里“切梯度”不是说某个 rank 永远见不到别人的梯度,而是 reduce-scatter 聚合后,每个 rank 只保留自己负责更新的那一片结果,其余梯度用完即释放。

ZeRO-3:参数也切。 参数、梯度和优化器状态全部只保留 $1/N$:

$$M_{\mathrm{Z3}}=\frac{16P}{N}$$

这时单卡模型状态随卡数线性下降,理论上终于可以使用整个数据并行组的聚合显存。代价是某层计算前必须 all-gather 出完整参数,计算后再 reshard。这里的“完整”通常只针对一个 FSDP unit 或若干层,而不是一次把全模型永远拼回显存;包装粒度如果选错,峰值仍然可能爆掉。

代入 $P=7.5\times10^9$、$N=64$:

  • DDP:$16P=120.00$ GB;
  • ZeRO-1:$4P+12P/64=31.41$ GB;
  • ZeRO-2:$2P+14P/64=16.64$ GB;
  • ZeRO-3:$16P/64=1.875$ GB。

这与 ZeRO 原论文表 1中的 120、31.4、16.6、1.88 GB 对上了。论文使用十进制 GB;如果监控工具显示 GiB,同一字节数会小约 7.4%,比较数字前先统一单位。

3.4 显存省了,通信为什么没有立刻爆炸

设梯度张量大小为 $G$ 字节。高带宽 ring all-reduce 通常拆成 reduce-scatter 和 all-gather。忽略延迟、拓扑和协议常数后,每个 rank 在两个阶段分别移动:

$$V_{\mathrm{RS}}=\frac{N-1}{N}G,\qquad V_{\mathrm{AG}}=\frac{N-1}{N}G$$

因此普通 DDP 的总通信量近似为:

$$V_{\mathrm{DDP}}=2\frac{N-1}{N}G\approx2G$$

ZeRO-1/2 把原来的 all-reduce 改排成“梯度 reduce-scatter + 更新后参数 all-gather”。在本文参数与梯度均为两字节的假设下,两者大小同为 $G$,理想总字节量仍约为 $2G$。若用 FP32 梯度而参数为 BF16,两项大小不同,不能直接沿用该字节比。区别在于数据留在哪里、何时释放,而不是凭空少传了一半。

ZeRO-3 还要在前向和反向各聚合一次参数,再做梯度 reduce-scatter。按 ZeRO 论文用参数元素数 $P$ 计量,普通数据并行约移动 $2P$ 个元素,Stage 3 约移动 $3P$ 个元素,所以是基线的 1.5 倍。这个 1.5 倍假设参数/梯度通信 dtype 相同并采用所述聚合与释放调度,是理想带宽模型,不是训练时间必然乘 1.5:通信能否和计算重叠、跨没跨节点、消息是否太碎,都会改变真实结果。

3.5 别把模型状态公式当成整机显存公式

实际峰值更接近:

$$M_{\mathrm{peak}}=M_{\mathrm{model\ states}}+M_{\mathrm{activations}}+M_{\mathrm{temporary}}+M_{\mathrm{runtime}}+M_{\mathrm{fragmentation}}$$

第一项就是前面推导的参数、梯度和优化器状态。激活由 batch、序列长度、空间分辨率、层数和 checkpoint 策略决定;临时项包含 all-gather bucket、梯度 bucket 和算子工作区;runtime 包含 CUDA context、通信库等;碎片则取决于张量生命周期和分配器状态。

ZeRO 原论文给过一个很有提醒意义的例子:15 亿参数 GPT-2 的模型状态至少 24 GB;序列长度 1024、batch 32 时,激活约 60 GB,即使用激活重计算降到约 8 GB,临时 FP32 扁平缓冲还可能再占 6 GB。只看 $16P$ 判断“32 GB 正好能放下”会直接翻车。

模型状态随数据并行卡数的变化

模型状态随数据并行卡数的变化。固定 FP16 参数/梯度与 FP32 主权重及 Adam 状态,纵轴为十进制 GB;不含激活、临时聚合缓冲和运行时。

04. 代码实现

本节有两个完整脚本,均只依赖 Python 标准库。第一个复现论文显存表;第二个把 DDP all-reduce 与 ZeRO-2 的 reduce-scatter + all-gather 拆开,验证它们做出同一次 Adam 更新。完整代码在文末附录。

4.1 用公式生成显存账本

zero_memory_ledger.py 把每个阶段写成“每卡复制多少字节 + 分片多少字节”:

STAGES = (
    Stage("DDP", 16.0, 0.0, 1.0),
    Stage("ZeRO-1", 4.0, 12.0, 1.0),
    Stage("ZeRO-2", 2.0, 14.0, 1.0),
    Stage("ZeRO-3", 0.0, 16.0, 1.5),
)

def bytes_per_parameter(self, world_size: int) -> float:
    return self.replicated_bytes + self.sharded_bytes / world_size

replicated_bytes 是每个 rank 必须完整保留的部分,sharded_bytes 是可以除以 $N$ 的部分。默认参数就是论文的 7.5B/64 卡案例。实际运行:

model=7.5B, world_size=64
assumption: fp16 params 2B + fp16 grads 2B + fp32 master/m/v 12B = 16 bytes/parameter
stage    bytes/param    model-state GB    vs DDP    comm/step GB
DDP         16.0000            120.00      1.00x           30.00
ZeRO-1       4.1875             31.41      3.82x           30.00
ZeRO-2       2.2188             16.64      7.21x           30.00
ZeRO-3       0.2500              1.88     64.00x           45.00
note: activation, temporary buffers and fragmentation are excluded

为什么 64 卡的 ZeRO-1 只有 3.82 倍而不是宣传里的 4 倍?因为 4 倍是 $N\to\infty$ 的上限,有限卡数下 $12P/N$ 还没有消失。ZeRO-2 的 7.21 倍同理。ZeRO-3 没有复制项,所以恰好获得 64 倍。

再换成 1.5B/8 卡:

model=1.5B, world_size=8
assumption: fp16 params 2B + fp16 grads 2B + fp32 master/m/v 12B = 16 bytes/parameter
stage    bytes/param    model-state GB    vs DDP    comm/step GB
DDP         16.0000             24.00      1.00x            6.00
ZeRO-1       5.5000              8.25      2.91x            6.00
ZeRO-2       3.7500              5.62      4.27x            6.00
ZeRO-3       2.0000              3.00      8.00x            9.00
note: activation, temporary buffers and fragmentation are excluded

这组数字揭示一个选型习惯:模型状态只差几 GB 时,ZeRO-1/2 往往已经够用;不必为了追求最低常驻显存直接承担 Stage 3 的参数聚合。

4.2 ZeRO 为什么没有改掉优化结果

zero_update_simulator.py 构造 4 个 rank、8 个参数。每个 rank 看见不同数据,因此产生不同局部梯度。DDP 路径先得到完整平均梯度,再完整执行 Adam;ZeRO-2 路径只把平均梯度对应的两元素分片交给各 rank,每个 rank 维护自己的 $m,v$,更新后再把四个参数片聚合起来。

核心区别只有分片发生的位置:

reduced_full_grad = average_columns(local_grads)
ddp_params, ddp_m, ddp_v = adam_first_step(PARAMS, reduced_full_grad)

param_shards = shard(PARAMS, WORLD_SIZE)
grad_shards = shard(reduced_full_grad, WORLD_SIZE)
updated_shards = []
for params_for_rank, grads_for_rank in zip(param_shards, grad_shards):
    updated, local_m, local_v = adam_first_step(params_for_rank, grads_for_rank)
    updated_shards.append(updated)
zero_params = [value for part in updated_shards for value in part]

真实运行输出:

world_size=4, parameter_count=8, shard_width=2
local gradient matrix shape=(4, 8)
reduce-scatter result shape=(4, 2)
averaged gradient: [0.02, 0.045, 0.07, 0.095, 0.12, 0.145, 0.17, 0.195]
DDP updated params: [0.099, -0.201, 0.299, -0.401, 0.499, -0.601, 0.699, -0.801]
ZeRO-2 gathered params: [0.099, -0.201, 0.299, -0.401, 0.499, -0.601, 0.699, -0.801]
max_abs_diff=0.000000000000
Adam m/v scalars per rank: DDP=16, ZeRO-2=4, reduction=4.00x

max_abs_diff=0 是本文最重要的代码结果:同一份平均梯度、同一优化器规则下,把参数更新分给不同 rank 并不会改变更新后的模型。每卡 Adam 的 $m,v$ 元素数则从 16 降到 4,正好是 4 倍。真实系统还需要 bucket、异步通信、混合精度和异常恢复,但数学骨架就是这几十行。

05. 工业级实现对照

最小实现为了看懂“切什么”,生产实现要解决的则是“什么时候切、什么时候聚合、怎么让网络传输藏在计算后面”。下面以 2026-09-05 可见的官方实现为准。

5.1 PyTorch DDP:复制参数,反向时同步梯度

PyTorch 的 DistributedDataParallel 采用一进程一卡。它不会自动切分输入;应用通常用 DistributedSampler 保证不同 rank 读取不同样本。模型参数在每卡完整复制,autograd hook 在反向过程中把就绪梯度装进 bucket,并尽早发起 all-reduce,以便通信和后续层反向计算重叠。

这解释了两个常见现象:第一,DDP 通常比单进程 DataParallel 快,因为没有主卡收集和 Python 线程瓶颈;第二,DDP 加卡能缩短时间,却不会降低模型状态的单卡显存。它首先是一种吞吐扩展方案,不是大模型装载方案。

5.2 ZeroRedundancyOptimizer:最小改动的 Stage 1 思路

PyTorch 的 ZeroRedundancyOptimizer 可以和 DDP 组合。每个 rank 只为大约 $1/N$ 的参数维护本地优化器状态,更新自己负责的参数后广播结果,让所有 DDP 副本重新一致。这对应 ZeRO-1 的核心思想;参数仍完整复制,所以不能把它当成 FSDP 的替代品。

官方文档还标记该 API 为 experimental,并提示启用 overlap_with_ddp=True 时,最初若干迭代可能因为梯度 bucket 尚未稳定而不做参数更新。工程上不能只看“少了多少 GB”,还要读清 API 的更新时序、checkpoint 聚合和版本状态。

5.3 DeepSpeed ZeRO:三个阶段直接写进配置

DeepSpeed 官方 ZeRO 教程把三个阶段定义得很直接:Stage 1 切优化器状态,Stage 2 再切梯度,Stage 3 再切参数。一个典型配置是:

{
  "zero_optimization": {
    "stage": 2,
    "contiguous_gradients": true,
    "overlap_comm": true,
    "reduce_bucket_size": 500000000
  }
}

contiguous_gradients 针对碎片与通信连续性,overlap_comm 尝试把通信藏进计算,reduce_bucket_size 在“消息大到能吃满带宽”和“临时 bucket 不要撑爆显存”之间取舍。这些开关没有脱离前面的公式,只是在控制公式之外的临时项和时间轴。

如果 GPU 显存仍不够,ZeRO-Offload 可以把优化器状态与计算移到 CPU,ZeRO-Infinity 进一步使用 CPU 和 NVMe。但 官方 ZeRO-Offload 教程 展示的 10B 单 V100 案例并不意味着 offload 免费:PCIe 传输、CPU Adam 吞吐、NUMA 与磁盘带宽会变成新的瓶颈。

5.4 PyTorch FSDP:Stage 2/3 的框架化实现

知识树里的代码锚点 torch/distributed/fsdp/fully_sharded_data_parallel.py#FullyShardedDataParallel 在当前 PyTorch main 仍存在。FSDP1 的 ShardingStrategy 可以这样理解:

  • NO_SHARD:参数、梯度、优化器状态都复制,行为接近 DDP;
  • SHARD_GRAD_OP:梯度和优化器状态分片,参数在计算窗口内保持完整,接近 ZeRO-2;
  • FULL_SHARD:参数、梯度和优化器状态全分片,接近 ZeRO-3;
  • HYBRID_SHARD:节点内全分片、节点间复制,减少低带宽跨节点通信。

不过当前 PyTorch FSDP2 教程 已经明确建议迁移到 fully_shard:FSDP2 以 DTensor 做逐参数分片,不再依赖 FSDP1 的扁平参数;reshard_after_forward=True 对应 FULL_SHARD,设为 False 则更像 SHARD_GRAD_OP。文章或配置里只写“用了 FSDP”已经不够,必须同时说明代际和策略。

5.5 一层参数在 FSDP 中的生命周期

以 full shard 为例,一个 FSDP unit 在一次训练步里大致经历:

  1. 常驻状态只有本 rank 的参数分片;
  2. pre-forward hook 发起 all-gather,临时物化完整参数;
  3. 执行该 unit 的前向;
  4. 释放完整参数,恢复分片;
  5. 反向前再次 all-gather 参数;
  6. 计算梯度后执行 reduce-scatter,每个 rank 只留下本地梯度片;
  7. 本地优化器只更新本 rank 的参数与状态分片。

真正决定峰值的不是“最终只存 $1/N$”,而是同一时刻有多少 unit 正在 all-gather、预取队列有多深、最大 unit 有多大。如果把整个模型只包成一个巨型 unit,那么计算前仍要暂时物化全模型;如果把每个很小的算子都单独包起来,又会制造大量小消息,延迟和调度开销反而吞掉吞吐。Transformer 常按 block 包装,就是在这两端之间折中。

FSDP2 的公开契约仍是同一个时间逻辑:前向/反向前由 hook unshard,之后 reshard,梯度用 reduce-scatter 汇聚。实现细节从 flat parameter 变成 DTensor,不代表 ZeRO 的“按需物化”思想变了。

5.6 一个实用的选择顺序

先测量,再逐级加复杂度:

  1. 模型状态能放下,只想提吞吐:先用 DDP;
  2. Adam 状态是主要缺口:DDP + ZeRO-1/ZeroRedundancyOptimizer;
  3. 梯度也造成明显压力:ZeRO-2 或 FSDP 的 shard-grad 策略;
  4. 单卡连参数副本都放不下:ZeRO-3/FSDP full shard;
  5. 模型状态已经很小但长视频仍 OOM:处理激活重计算、micro-batch、FlashAttention 或序列并行,而不是继续折腾 ZeRO stage;
  6. 网络太慢:优先让分片组留在节点内,再用张量/流水线或混合分片扩到节点间。

06. 代价与边界

ZeRO 不减少总计算量。 同一个 batch 的前向、反向和 Adam 更新并没有少。它让每个 rank 少存状态,并通过通信在需要时恢复视图。若模型原本就能舒适放下,Stage 3 很可能只是增加通信和 hook 调度,吞吐反而下降。

通信字节相同,不代表时间相同。 ZeRO-2 与 DDP 在论文模型里都是约 $2G$,但一次大 all-reduce 和许多 layer-wise reduce-scatter/all-gather 的延迟特征不同。NVLink 节点内、InfiniBand 节点间、普通以太网的最优 bucket 大小不会一样;跨节点带宽不足时,1.5 倍 Stage 3 通信尤其明显。

平均显存很低,峰值仍可能 OOM。 参数 all-gather、预取、梯度归约、算子 workspace 可能同时在场。只看稳定阶段的 memory_allocated 会漏掉峰值;要记录 max_memory_allocated,并逐步调 wrap 粒度、prefetch 和 bucket。

ZeRO 解决不了激活随序列长度增长。 视频 DiT 中,时间、宽、高一起扩张,token 数可能成倍增加,注意力和 MLP 保存的激活随之增长。参数切到 1 GB 后仍然爆显存,并不说明 ZeRO 失效,而是瓶颈已经从模型状态转移到激活。

checkpoint 变复杂。 每个 rank 手里只有一片状态,保存时要选择 full、sharded 或 local state dict。把完整 checkpoint 聚合到 rank 0 可能让 CPU 内存瞬间成为瓶颈;只保存分片又要求恢复时正确处理 world size 和布局。训练能跑并不等于容灾链路可用,上线前至少做一次“保存—退出—换步数恢复”的演练。

全局操作必须认识分片。 梯度范数裁剪、参数检查、EMA、冻结部分参数、权重共享都不能默认“当前 rank 能看到完整张量”。例如 FSDP 提供自己的 clip_grad_norm_,就是因为全局范数需要跨分片归约。绕开框架直接遍历本地 .grad,算到的只是局部值。

offload 是拿带宽换容量。 CPU 内存比显存大,NVMe 又比 CPU 内存大,但层级越远,带宽越低、延迟越高。小模型或计算密度不够的模型会被数据搬运压垮;只有“不 offload 根本跑不了”或计算足以覆盖传输时,这个交换才划算。

大 batch 会改变优化问题。 增加数据并行度时,如果每卡 batch 不变,全局 batch 会随 $N$ 增长。ZeRO 保证同一全局 batch 下更新等价,却不保证扩大 batch 后收敛曲线不变。学习率、warmup、梯度累计和数据采样仍要一起调整。

什么时候不该用 ZeRO-3:模型与激活在单卡尚有充足余量、网络较慢、模型由大量极小模块组成且难以形成高效 bucket,或者你更看重最低延迟和调试简单性。此时 DDP 或 ZeRO-1/2 往往是更好的工程答案。

07. 经典论文脉络

这条路线不是“突然发明一种分布式训练”,而是一步步把数据并行里的冗余拆掉:

  • Horovod(arXiv:1802.05799,2018)把 ring all-reduce 做成易接入的训练抽象,奠定了现代同步数据并行“各算各的、梯度集体归约”的工程基线;它解决吞吐扩展,但每个 worker 仍保存完整模型状态。
  • ZeRO(arXiv:1910.02054,2019)指出数据并行浪费的不是总显存,而是参数、梯度和优化器状态的重复副本;三阶段切分在保持数据并行计算粒度的同时,把模型状态显存从 $16P$ 推到 $16P/N$。
  • ZeRO-Offload(arXiv:2101.06840,2021)把优化器状态与计算搬到 CPU,并针对 CPU Adam 优化,让单张 32 GB V100 训练 10B 模型成为论文展示案例;留下的新问题是主机带宽和异构调度。
  • ZeRO-Infinity(arXiv:2104.07857,2021)继续把 CPU 与 NVMe 纳入内存层级,用带宽感知的分区和预取突破 GPU 内存墙;容量继续扩大,但 I/O 调度成为系统的核心。
  • PyTorch FSDP(arXiv:2304.11277,2023)总结了把 fully sharded data parallelism 纳入 PyTorch eager 训练栈的实践,让 ZeRO-3 类思想不再只属于单一外部训练引擎,并推动后续 FSDP2 的逐参数分片。

主线可以压缩成一句话:all-reduce 证明“计算可以复制、结果可以同步”,ZeRO 进一步问“既然结果会同步,状态为什么还要复制”。 后续 Offload、Infinity、FSDP 都是在回答状态放在哪、何时出现、以什么粒度搬运。

08. 常见误解

“数据并行会把模型切到多张卡上。” 不会。普通 DDP 是每卡一个完整模型,只切数据。能把吞吐从一张卡扩到八张,不代表能装下超过单卡容量的参数。

“ZeRO-3 后每卡永远只有 $1/N$ 参数。” 常驻状态是 $1/N$,计算某个 FSDP unit 前仍需临时 all-gather 完整参数。忽略这个瞬时窗口,就会得到理论显存能放、实际 forward 前仍 OOM 的配置。

“ZeRO-2 比 DDP 少传一半梯度,所以一定更快。” ZeRO-2 用 reduce-scatter 只留下梯度片,但更新后还要 all-gather 参数片。按论文的理想字节模型,总通信量与 DDP all-reduce 相同;快慢来自重叠、bucket、拓扑与实现,而非简单少一半。

“用了 ZeRO 就不用 activation checkpointing。” 两者切的是不同账本。ZeRO 处理模型状态冗余;activation checkpointing 用额外重算减少前向激活保存。长视频训练经常需要两者同时用。

“参数是 BF16,所以 Adam 也只占 2 字节。” 常见混合精度训练会保留 FP32 主参数与 FP32 的 $m,v$,优化器部分仍是 12 字节/参数。不同优化器、8-bit optimizer 或纯 BF16 更新会改变常数,因此使用 $16P$ 前必须先列出真实 dtype 与状态。

“64 卡 ZeRO-1 就一定省 4 倍。” 4 倍是 $N$ 很大时的渐近上限。64 卡的精确值是 $16/(4+12/64)=3.82$ 倍;8 卡只有 2.91 倍。宣传中的“up to”不能替代自己的显存账本。

“只要 loss 一样,所有数据并行更新都严格一致。” 浮点归约顺序会改变末位误差,随机数、dropout、数据尾 batch 和非确定性 kernel 也会影响复现。数学上等价不等于 bitwise identical;本文模拟得到 0 误差,是因为使用同一确定性顺序与双精度标量。

09. 动手验证

下面五个实验都能在没有 GPU 的机器上完成。两个脚本全文在附录,复制保存后直接用 Python 3 运行。

实验一:复现原论文表格。

python zero_memory_ledger.py

预期看到 7.5B/64 卡下 DDP、ZeRO-1/2/3 分别为 120.00、31.41、16.64、1.88 GB。如果结果差约 7%,检查自己是不是把 GiB 和十进制 GB 混用了。

实验二:观察有限卡数离理论上限有多远。

python zero_memory_ledger.py --params-b 1.5 --world-size 8

预期 ZeRO-1 只省 2.91 倍,ZeRO-2 只省 4.27 倍,ZeRO-3 才精确省 8 倍。再把 world-size 改成 2、4、16、64,观察 Stage 1 向 4 倍、Stage 2 向 8 倍收敛。

实验三:验证更新等价性。

python zero_update_simulator.py

预期 DDP 与 ZeRO-2 的八个新参数完全相同,max_abs_diff=0;每 rank 的 Adam $m,v$ 标量数从 16 降为 4。把 local_gradient() 的公式改掉,只要两条路径仍使用同一平均梯度,最终参数就应继续一致。

实验四:故意制造尾 batch 权重错误。 在模拟器里让最后一个 rank 只代表一个样本、其他 rank 各代表两个样本,然后比较“先对每个 rank 求平均再除以 4”和“按总样本数加权平均”。预期两个梯度不同。这能解释为什么真实训练要使用 DistributedSampler、drop_last 或正确的样本权重,而不能只相信 all-reduce。

实验五:给自己的模型补全账本。 从训练日志记录参数量、每种状态 dtype、峰值激活与临时 buffer。先用脚本算模型状态理论值,再和框架峰值相减。如果差值随序列长度或分辨率快速增长,瓶颈是激活;如果几乎不随输入变化,才优先继续切模型状态。预期这一步比直接试三个 ZeRO stage 更快找到真正的 OOM 原因。

10. 延伸阅读

读完这篇可以沿分布式训练知识树继续:

  • 混合精度与数值稳定性:本文的 16 字节常数建立在混合精度 Adam 上;理解 FP16、BF16、FP8 的动态范围后,才能判断哪些状态可以继续降精度。
  • 张量并行与流水线并行:当单个算子或单层连临时 all-gather 都放不下时,需要切计算本身;下一步要算 TP 通信量和 PP 气泡率。
  • 序列并行与 Ring Attention:模型状态已经切完,长视频仍被激活卡住时,才轮到沿序列维度切注意力计算。

本文的核心资料均来自原始论文与官方实现:ZeRO 的显存和通信公式以 ZeRO 原论文为准;三阶段配置参考 DeepSpeed ZeRO 官方教程;DDP、优化器分片和 FSDP 策略分别核对了 PyTorch DDP 文档、ZeroRedundancyOptimizer 文档 与 FSDP 文档。代码路径和 API 会随上游重构,工业实现部分注明的日期就是版本边界。

附录:完整代码

09 节用到的脚本全文如下(zero_memory_ledger.py、zero_update_simulator.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

zero_memory_ledger.py

#!/usr/bin/env python3
"""复现 ZeRO 论文中的混合精度 Adam 模型状态显存账本。

这里只计算参数、梯度和优化器状态,不包含激活、临时通信缓冲区、
CUDA context 与内存碎片。单位同时使用十进制 GB,便于和论文表格对照。
"""

from __future__ import annotations

import argparse
from dataclasses import dataclass


@dataclass(frozen=True)
class Stage:
    name: str
    replicated_bytes: float
    sharded_bytes: float
    communication_multiple: float

    def bytes_per_parameter(self, world_size: int) -> float:
        return self.replicated_bytes + self.sharded_bytes / world_size


STAGES = (
    Stage("DDP", 16.0, 0.0, 1.0),
    Stage("ZeRO-1", 4.0, 12.0, 1.0),
    Stage("ZeRO-2", 2.0, 14.0, 1.0),
    Stage("ZeRO-3", 0.0, 16.0, 1.5),
)


def gb(byte_count: float) -> float:
    """转为十进制 GB;ZeRO 原论文的 120 GB 使用这个口径。"""
    return byte_count / 1_000_000_000


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--params-b", type=float, default=7.5,
                        help="参数量,单位十亿;默认复现论文的 7.5B 示例")
    parser.add_argument("--world-size", type=int, default=64,
                        help="数据并行进程数")
    args = parser.parse_args()
    if args.params_b <= 0 or args.world_size <= 0:
        raise ValueError("params-b 与 world-size 必须为正数")

    params = args.params_b * 1_000_000_000
    gradient_bytes = params * 2  # fp16/bf16 梯度
    baseline_comm = 2 * gradient_bytes  # ring all-reduce: reduce-scatter + all-gather

    print(f"model={args.params_b:g}B, world_size={args.world_size}")
    print("assumption: fp16 params 2B + fp16 grads 2B + "
          "fp32 master/m/v 12B = 16 bytes/parameter")
    print("stage    bytes/param    model-state GB    vs DDP    comm/step GB")
    ddp_bytes = STAGES[0].bytes_per_parameter(args.world_size)
    for stage in STAGES:
        bpp = stage.bytes_per_parameter(args.world_size)
        state_gb = gb(params * bpp)
        saving = ddp_bytes / bpp
        comm_gb = gb(baseline_comm * stage.communication_multiple)
        print(f"{stage.name:<8} {bpp:>10.4f} {state_gb:>17.2f} "
              f"{saving:>9.2f}x {comm_gb:>15.2f}")

    print("note: activation, temporary buffers and fragmentation are excluded")


if __name__ == "__main__":
    main()

zero_update_simulator.py

#!/usr/bin/env python3
"""用单进程模拟 DDP 与 ZeRO-2 的一次 Adam 更新。

脚本不依赖 GPU、PyTorch 或分布式运行时。四个“rank”先各自产生一份局部
梯度,再比较两条路径:DDP 对完整梯度做 all-reduce;ZeRO-2 对梯度做
reduce-scatter、每个 rank 只更新自己的参数片,最后 all-gather 参数。
"""

from __future__ import annotations

import math


WORLD_SIZE = 4
PARAMS = [0.10, -0.20, 0.30, -0.40, 0.50, -0.60, 0.70, -0.80]
LEARNING_RATE = 1e-3
BETA1 = 0.9
BETA2 = 0.999
EPSILON = 1e-8


def local_gradient(rank: int, width: int) -> list[float]:
    """构造确定性的局部梯度,模拟每个 rank 看见不同数据。"""
    return [((rank + 1) * (index + 2) - 3) / 100.0 for index in range(width)]


def average_columns(rows: list[list[float]]) -> list[float]:
    return [sum(column) / len(rows) for column in zip(*rows)]


def shard(vector: list[float], world_size: int) -> list[list[float]]:
    if len(vector) % world_size:
        raise ValueError("为了让示例清楚,参数量必须能被 world_size 整除")
    width = len(vector) // world_size
    return [vector[rank * width:(rank + 1) * width]
            for rank in range(world_size)]


def adam_first_step(params: list[float], grads: list[float]) -> tuple[list[float], list[float], list[float]]:
    """执行 Adam 的第 1 步,并返回新参数、m、v。"""
    new_params: list[float] = []
    m_state: list[float] = []
    v_state: list[float] = []
    for value, grad in zip(params, grads):
        m = (1.0 - BETA1) * grad
        v = (1.0 - BETA2) * grad * grad
        m_hat = m / (1.0 - BETA1)
        v_hat = v / (1.0 - BETA2)
        updated = value - LEARNING_RATE * m_hat / (math.sqrt(v_hat) + EPSILON)
        new_params.append(updated)
        m_state.append(m)
        v_state.append(v)
    return new_params, m_state, v_state


def main() -> None:
    local_grads = [local_gradient(rank, len(PARAMS))
                   for rank in range(WORLD_SIZE)]

    # DDP:每个 rank 经 all-reduce 得到同一份完整平均梯度,并完整更新参数。
    reduced_full_grad = average_columns(local_grads)
    ddp_params, ddp_m, ddp_v = adam_first_step(PARAMS, reduced_full_grad)

    # ZeRO-2:reduce-scatter 的结果等价于先平均,再只保留所属分片。
    param_shards = shard(PARAMS, WORLD_SIZE)
    grad_shards = shard(reduced_full_grad, WORLD_SIZE)
    updated_shards: list[list[float]] = []
    state_sizes: list[int] = []
    for params_for_rank, grads_for_rank in zip(param_shards, grad_shards):
        updated, local_m, local_v = adam_first_step(params_for_rank, grads_for_rank)
        updated_shards.append(updated)
        state_sizes.append(len(local_m) + len(local_v))

    # all-gather 后每个 rank 都能拿到同一份新参数;这里拼接一次代表该结果。
    zero_params = [value for part in updated_shards for value in part]
    max_abs_diff = max(abs(a - b) for a, b in zip(ddp_params, zero_params))

    shard_width = len(PARAMS) // WORLD_SIZE
    print(f"world_size={WORLD_SIZE}, parameter_count={len(PARAMS)}, "
          f"shard_width={shard_width}")
    print(f"local gradient matrix shape=({WORLD_SIZE}, {len(PARAMS)})")
    print(f"reduce-scatter result shape=({WORLD_SIZE}, {shard_width})")
    print("averaged gradient:", [round(x, 4) for x in reduced_full_grad])
    print("DDP updated params:", [round(x, 6) for x in ddp_params])
    print("ZeRO-2 gathered params:", [round(x, 6) for x in zero_params])
    print(f"max_abs_diff={max_abs_diff:.12f}")
    print(f"Adam m/v scalars per rank: DDP={len(ddp_m) + len(ddp_v)}, "
          f"ZeRO-2={state_sizes[0]}, reduction={WORLD_SIZE:.2f}x")


if __name__ == "__main__":
    main()

make_figures.py

"""Regenerate this article's deterministic teaching figure (numpy + matplotlib)."""
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
OUT=Path(__file__).resolve().parents[1]/"figures"
OUT.mkdir(exist_ok=True)
plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False})
n=np.array([1,2,4,8,16,32,64]); p=7.5e9
fig,ax=plt.subplots(figsize=(8,4.8))
for name,y in [("DDP",np.full(n.shape,16.)),("ZeRO-1",4+12/n),("ZeRO-2",2+14/n),("ZeRO-3",16/n)]:
    ax.plot(n,y*p/1e9,"o-",label=name)
ax.set(xscale="log",yscale="log",xlabel="Data-parallel ranks",ylabel="Model states per rank (GB)",title="7.5B model: FP16 weights/grads + FP32 master/m/v")
ax.legend();ax.grid(alpha=.25);fig.tight_layout()
fig.savefig(OUT/'zero_states.png',dpi=170)
plt.close(fig)
print(OUT/'zero_states.png')


更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」

每日更新 · 论文精选 · 深度解读 · 技术脉络

微信搜索 人工智能炼丹君 或扫描下方二维码关注

扫码关注「人工智能炼丹君」

0

评论 (0)

取消
粤ICP备2021042327号