本材料说明:全文数字均复述自论文原文(arXiv v3, 2020-05-13),未做外部验证,出处以 §小节号 标注。【论文声称】=作者观点/叙事框架;【实验支持】=文中实验数据支撑;【解读者推断】=精读者的推理,论文未明说。
训练时真正吃显存的不是参数本身,而是被复制到每张卡上的优化器状态、梯度和参数;ZeRO 的做法是把这三类"模型状态"分区(partition)而非复制(replicate),在通信量几乎不增加的前提下,把每卡模型状态内存从 16Ψ 字节压到 4Ψ → 2Ψ → 16Ψ/Nd 字节(Ψ 为参数量,Nd 为数据并行度),从而宣称用 1024 张 GPU 就能装下 1 万亿参数模型的模型状态(§1)。
| 摘要短语(§Abstract) | 对应论文章节 | 对应本精读章 |
|---|---|---|
| "eliminates memory redundancies in data- and model-parallel training" | §1、§4–§5(ZeRO-DP 三级分区) | 第三章 |
| "retaining low communication volume" | §7(ZeRO-DP 通信分析)、§8(ZeRO-R 通信分析) | 第三章、第四章 |
| "scale beyond 1 Trillion parameters using today's hardware" | §5.4、Table 1/2、§9 | 第三章、第五章 |
| "residual states"(激活/临时缓冲/碎片)优化 | §3.2、§6(ZeRO-R:Pa/CB/MD) | 第四章 |
| "100B+ on 400 GPUs, 15 Petaflops, 8x model size, 10x speedup" | §10(实现与评测,ZeRO-100B) | 第五章 |
| "largest language model (17B) with record breaking accuracy" | §10.6(Turing-NLG) | 第六章(批判性阅读) |
必读主线:§3(内存账目)→ §5(三级分区)→ §7(通信分析)→ Table 1。这条线是全文的骨架,跳过任何一环都会在后续章节迷失。可跳读支线:§2 相关工作、§6 ZeRO-R 细节、§10 评测细节。跳过 §6 的代价是看不懂 §10.5 的 C1–C5 配置消融;跳过 §2 无碍主线。
来源:论文 §3.1–§3.2。
GPT-2 有 1.5B 参数,fp16 权重只占 3GB,但论文指出它无法在单张 32GB GPU 上用 TensorFlow 或 PyTorch 训练(§3 开篇)。缺的 21GB 去哪了?答案是三类隐藏开销:
16Ψ 字节(§3.1)。GPT-2 即需 ≥24GB——是权重的 8 倍。| 符号 | 它是什么 | 直觉 |
|---|---|---|
| Ψ | 模型参数个数 | 一切账目的基数 |
| 2Ψ | fp16 参数占的字节 | 每参数 2 字节 |
| 2Ψ | fp16 梯度占的字节 | 反向传播的产物 |
| K=12 | 优化器状态乘子:fp32参数4Ψ + 动量4Ψ + 方差4Ψ | Adam 的"记忆"比参数本身贵 6 倍 |
手算验证:GPT-2 Ψ=1.5×109,M = 16 × 1.5e9 = 2.4×1010 B ≈ 22.4 GiB,论文写"至少 24 GB"(按十进制 GB),对得上(§3.1)。数字对账:§3 开篇说权重 3GB(=2×1.5e9 十进制),24GB/3GB=8,恰为 16字节/2字节=8 倍,两处自洽 ✓。
常见误读:① "参数多大就占多少显存"——错,训练态下优化器状态才是大头(KΨ=12Ψ > 参数的 4Ψ);② "混合精度省一半显存"——只省了参数和激活那部分,fp32 主权重+动量+方差一点没少;③ 论文脚注给出的激活公式 12×hidden×batch×seq×layers 只适用于 GPT-2 类架构的粗估,不能直接套到所有 Transformer 变体上。
一句话记住本节:训练显存的大头不是参数而是"围绕参数的状态",混合精度 Adam 是参数字节数的 8 倍。
学完本章你应能:①默写三个阶段的内存公式并解释每一项来源;②手算任意 (Ψ, Nd) 组合下的每卡模型状态内存;③说出为什么 Pos+g 不增加总通信量、而 Pp 恰好多 50%;④指出"ZeRO 改变了数学语义吗"这个问题的答案(不变——它只是重新分配了存储位置和通信时机)。
| 阶段 | 切什么 | 机制 | 通信代价 |
|---|---|---|---|
| Pos | 优化器状态(12Ψ) | 第 i 个进程只存并更新第 i 份优化器状态,步末 all-gather 同步新参数 | 与基线 DP 相同 |
| +Pg | 梯度(2Ψ) | reduce-scatter 把每层梯度归约到负责该参数分区的进程,用完即释放;bucket 化以重叠通信计算 | 相同(见3.4) |
| +Pp | 参数(2Ψ) | forward/backward 到哪层就 broadcast/all-gather 哪层参数,用完即弃,流水化调度 | 总计 3Ψ = 1.5x(见3.4) |
取 Figure 1 的例子:Ψ=7.5B,Nd=64,K=12。
再看 1T 模型行(Table 1):16000GB ÷ 1024 = 15.6GB,与正文"16TB/1024=16GB 装得进 32GB V100"(§1)口径一致(16TB 为约数)✓。对账发现:注意 Table 1 里 128B 模型在 Nd=256 时 Pos+g+p 已只需 8GB,说明"能装下"的门槛主要由 Nd 决定——这正是"超线性扩展"(§10.3)的根源:加卡不仅加算力,还直接加显存。
| 符号 | 它是什么 | 直觉 |
|---|---|---|
| Ψ(元素数) | 一次全量参数大小的数据搬运量 | all-gather 搬 Ψ、reduce-scatter 也搬 Ψ |
| Nd | 数据并行度 | 分片份数;Pp 的两次 all-gather 各摊销回 Ψ |
关键洞察(【解读者推断】,由 §7 推导过程支持):基线的 all-reduce 本来就等价于一次 reduce-scatter 加一次 all-gather,所以 Pos+g 只是"把本来就存在的两个操作拆开、插进不同时机",总量当然不变;Pp 多出来的是 backward 结束后为下一轮 forward 准备参数的那次额外 all-gather。
三级分区各自缓解的是显存压力;新增的压力分别是:Pos 要求步末同步等待(延迟);Pg 引入 bucket 管理复杂度;Pp 把通信从"每步一次大批量"变成"全程细粒度流水",对调度器要求高,且小消息带宽利用率低——这就是 DeepSeek 后续工程里反复出现的债。给后文埋的债:Pos+g+p 让每卡只剩 1/Nd 的状态,于是"加卡反而更快"的超线性区出现(§10.3),但这也意味着性能对集群规模和拓扑变得敏感——这是 §18 篇 MegaScale 要偿还的债。
失效模式先行:模型并行(MP)虽然切了参数,但每个 MP 进程都需要完整的激活来算自己那一竖条——激活被隐式复制了 Nm 份。ZeRO 的做法:forward 算完一层就把输入激活分区存放,backward 需要时再 all-gather 重建成完整副本(§6.1)。手算:100B 模型、MP=16、batch 32、seq 1024,每层存一份激活检查点需 ~33GB/卡;Pa 降到 ~2GB/卡(33/16),还可进一步 offload 到 CPU(Pa+cpu),激活占用趋近于零(§6.1)。
通信账:Megatron 每 transformer block 本有 6 次 all-reduce,总通信 12×seq×hidden;Pa 每次 backward 重算前多一次 all-gather(seq×hidden),增量 <10%(§8)。类比失效处:像"图书馆分馆藏书"没错,但这本书每次被借阅都要从各馆凑齐全本(all-gather),频繁借阅的小模型会被凑书成本拖垮——所以论文强调只对算术强度 ≥10K 的大模型启用(§4.2.1)。
CB(常量大小缓冲):把与模型大小成正比的扁平临时缓冲改成固定大小——3B 模型的 fp32 缓冲要 12GB,改成常量大缓冲后既保住带宽又不随模型膨胀(§6.2)。MD(即时碎片整理):长短生命周期张量交错是碎片根源(checkpoint 长寿命 vs 重计算激活短寿命;参数梯度长寿命 vs 激活梯度短寿命),做法是为这两类各预分配连续大块,产出即拷入(§6.3)。
一句话记住本节:ZeRO-R 管"剩下的"显存——激活靠分区(顺带救了MP的隐性复制),缓冲靠定长,碎片靠预分配。
闭卷自检清单:不看材料,我能说出——①Pa 解决的是 MP 的哪个隐性问题?②100B 模型激活检查点 33GB→2GB 依赖哪个除数?③MD 为什么把张量按生命周期分成两类?④CB 为什么不用"越大越好"的缓冲?
配置:实现的是 Pos+g + ZeRO-R(未含 Pp!),称 ZeRO-100B;硬件为 400 张 V100(25 个 DGX-2 节点),节点间 800 Gbps;基线为 PyTorch DDP(无MP)与 Megatron-LM 2019年9月开源版(有MP)(§10.1)。
核心结果(§10.2):最大跑到 170B 参数(SOTA Megatron 单独只能高效跑 ≤40B,8x 提升);8B–100B 模型平均持续吞吐 15 PetaFlops(>峰值 30%),单卡 >38 TFlops;相对基线最高 10x 提速。超线性扩展(§10.3):60B 模型从 64 卡到 400 卡,翻倍卡数性能翻倍以上。民主化(§10.4):无 MP 训 13B(128卡,>40 TFlops/卡),而 DDP 上限是 1.4B(<20 TFlops/卡)。
万亿之路的现实约束(§9):即便 ZeRO 让 1T 模型"装得下",也算不完——BERT-Large 在 1024 卡 DGX-2H 上 67 分钟训完;1T 模型每样本计算约为其 3000 倍,同硬件同效率也要 140 天起,实际上超一年;结论是需要 exa-FLOP 级系统。常见误读:"ZeRO 论文训练了万亿参数模型"——没有,1T 只是内存可行性分析,实际最大是 Turing-NLG 17B(Webtext-103 ppl 10.21,41.4 TFlops/GPU,§10.6)。
一句话记住本节:ZeRO-100B 用"半个ZeRO"(Pos+g+R)就把可训规模推到170B——剩下的一半(Pp)留给未来,也是后来 ZeRO-2/3 的伏笔。
论文的叙事轴是"模型大小×吞吐",但对以下成本着墨少:Pp 未实现未测(其 1.5x 通信只是纸面分析);Pa+cpu 在多数情况下降性能、仅在极端大模型下才划算(§10.5 自己承认 C5 在 60B 上反而更慢);开发成本上 ZeRO 易用性确实高,但 MD/Pa 的调度逻辑复杂度转移到了框架内部。