所属方向:分布式训练 | 难度:进阶 | 前置知识:无
关键词:数据并行、DDP、ZeRO、优化器状态切分、FSDP、显存占用估算
先算一笔会让很多人意外的账:训练一个 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 不改变模型的数学计算,而是按顺序切掉三种冗余副本:
把刚才的 75 亿参数模型放到 64 个数据并行 rank 上,三阶段的模型状态显存分别是 31.41 GB、16.64 GB 和 1.88 GB。最后一个数字正好是 $120/64$。这些数字不是宣传页上的模糊“最高节省多少”,而是本文后面会逐项推出来、再用代码复现的结果。
对视频生成尤其要补一句:视频 DiT 的 token 数很长,激活往往比模型状态还大。ZeRO 能拆掉模型状态冗余,但不会自动消灭激活。先分清是哪一类显存爆了,再决定用 ZeRO、激活重计算还是序列并行,比盲目把配置改成 stage 3 更重要。
三句话先建立框架:
如果只记一件事:DDP 是复制计算,ZeRO 是切分状态;ZeRO 没有改变梯度本身。
这里的 rank 可以先理解为一个独立训练进程,通常一个 rank 独占一张 GPU。$N$ 表示数据并行进程数,world_size 就是 $N$。每个 rank 处理本地 batch,但每一步结束后所有 rank 必须得到一致参数,否则下一步就不再是同一个模型。
设全局 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$。
设模型有 $P$ 个可训练参数,采用低精度前后向和 FP32 Adam 更新。每个参数对应:
后面三项统称优化器状态,共 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 假设严格对齐。
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$:
这与 ZeRO 原论文表 1中的 120、31.4、16.6、1.88 GB 对上了。论文使用十进制 GB;如果监控工具显示 GiB,同一字节数会小约 7.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:通信能否和计算重叠、跨没跨节点、消息是否太碎,都会改变真实结果。
实际峰值更接近:
$$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;不含激活、临时聚合缓冲和运行时。
本节有两个完整脚本,均只依赖 Python 标准库。第一个复现论文显存表;第二个把 DDP all-reduce 与 ZeRO-2 的 reduce-scatter + all-gather 拆开,验证它们做出同一次 Adam 更新。完整代码在文末附录。
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 的参数聚合。
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、异步通信、混合精度和异常恢复,但数学骨架就是这几十行。
最小实现为了看懂“切什么”,生产实现要解决的则是“什么时候切、什么时候聚合、怎么让网络传输藏在计算后面”。下面以 2026-09-05 可见的官方实现为准。
PyTorch 的 DistributedDataParallel 采用一进程一卡。它不会自动切分输入;应用通常用 DistributedSampler 保证不同 rank 读取不同样本。模型参数在每卡完整复制,autograd hook 在反向过程中把就绪梯度装进 bucket,并尽早发起 all-reduce,以便通信和后续层反向计算重叠。
这解释了两个常见现象:第一,DDP 通常比单进程 DataParallel 快,因为没有主卡收集和 Python 线程瓶颈;第二,DDP 加卡能缩短时间,却不会降低模型状态的单卡显存。它首先是一种吞吐扩展方案,不是大模型装载方案。
ZeroRedundancyOptimizer:最小改动的 Stage 1 思路PyTorch 的 ZeroRedundancyOptimizer 可以和 DDP 组合。每个 rank 只为大约 $1/N$ 的参数维护本地优化器状态,更新自己负责的参数后广播结果,让所有 DDP 副本重新一致。这对应 ZeRO-1 的核心思想;参数仍完整复制,所以不能把它当成 FSDP 的替代品。
官方文档还标记该 API 为 experimental,并提示启用 overlap_with_ddp=True 时,最初若干迭代可能因为梯度 bucket 尚未稳定而不做参数更新。工程上不能只看“少了多少 GB”,还要读清 API 的更新时序、checkpoint 聚合和版本状态。
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 与磁盘带宽会变成新的瓶颈。
知识树里的代码锚点 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”已经不够,必须同时说明代际和策略。
以 full shard 为例,一个 FSDP unit 在一次训练步里大致经历:
真正决定峰值的不是“最终只存 $1/N$”,而是同一时刻有多少 unit 正在 all-gather、预取队列有多深、最大 unit 有多大。如果把整个模型只包成一个巨型 unit,那么计算前仍要暂时物化全模型;如果把每个很小的算子都单独包起来,又会制造大量小消息,延迟和调度开销反而吞掉吞吐。Transformer 常按 block 包装,就是在这两端之间折中。
FSDP2 的公开契约仍是同一个时间逻辑:前向/反向前由 hook unshard,之后 reshard,梯度用 reduce-scatter 汇聚。实现细节从 flat parameter 变成 DTensor,不代表 ZeRO 的“按需物化”思想变了。
先测量,再逐级加复杂度:
ZeroRedundancyOptimizer;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 往往是更好的工程答案。
这条路线不是“突然发明一种分布式训练”,而是一步步把数据并行里的冗余拆掉:
主线可以压缩成一句话:all-reduce 证明“计算可以复制、结果可以同步”,ZeRO 进一步问“既然结果会同步,状态为什么还要复制”。 后续 Offload、Infinity、FSDP 都是在回答状态放在哪、何时出现、以什么粒度搬运。
“数据并行会把模型切到多张卡上。” 不会。普通 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 误差,是因为使用同一确定性顺序与双精度标量。
下面五个实验都能在没有 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 原因。
读完这篇可以沿分布式训练知识树继续:
本文的核心资料均来自原始论文与官方实现:ZeRO 的显存和通信公式以 ZeRO 原论文为准;三阶段配置参考 DeepSpeed ZeRO 官方教程;DDP、优化器分片和 FSDP 策略分别核对了 PyTorch DDP 文档、ZeroRedundancyOptimizer 文档 与 FSDP 文档。代码路径和 API 会随上游重构,工业实现部分注明的日期就是版本边界。
09 节用到的脚本全文如下(zero_memory_ledger.py、zero_update_simulator.py、make_figures.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()
#!/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()
"""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)