← 总目录 / 板块二 · Infra与数据的变迁
板块二 · Infra与数据的变迁

第13篇 · ZeRO

显存都去哪了?——用"分片"替代"复制",让数据并行也能装下万亿参数
ZeRO: Memory Optimizations Toward Training Trillion Parameter Models · Rajbhandari, Rasley, Ruwase, He(微软)· 2019 · arXiv:1910.02054

本材料说明:全文数字均复述自论文原文(arXiv v3, 2020-05-13),未做外部验证,出处以 §小节号 标注。【论文声称】=作者观点/叙事框架;【实验支持】=文中实验数据支撑;【解读者推断】=精读者的推理,论文未明说。

一、全局大图

1.1 一句话读懂本文

训练时真正吃显存的不是参数本身,而是被复制到每张卡上的优化器状态、梯度和参数;ZeRO 的做法是把这三类"模型状态"分区(partition)而非复制(replicate),在通信量几乎不增加的前提下,把每卡模型状态内存从 16Ψ 字节压到 4Ψ → 2Ψ → 16Ψ/Nd 字节(Ψ 为参数量,Nd 为数据并行度),从而宣称用 1024 张 GPU 就能装下 1 万亿参数模型的模型状态(§1)。

1.2 摘要—章节对照导航表

摘要短语(§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)第六章(批判性阅读)

1.3 主要贡献与证据强度预评

1.4 推荐阅读路线

必读主线:§3(内存账目)→ §5(三级分区)→ §7(通信分析)→ Table 1。这条线是全文的骨架,跳过任何一环都会在后续章节迷失。可跳读支线:§2 相关工作、§6 ZeRO-R 细节、§10 评测细节。跳过 §6 的代价是看不懂 §10.5 的 C1–C5 配置消融;跳过 §2 无碍主线。

二、逐章精读 · 内存去哪了(对应论文§3)

来源:论文 §3.1–§3.2。

2.1 失效模式先行:3GB 的模型为什么装不进 32GB 的卡?

GPT-2 有 1.5B 参数,fp16 权重只占 3GB,但论文指出它无法在单张 32GB GPU 上用 TensorFlow 或 PyTorch 训练(§3 开篇)。缺的 21GB 去哪了?答案是三类隐藏开销:

  1. 优化器状态:混合精度训练下,Adam 需要维护 fp32 参数副本 + fp32 动量 + fp32 方差,共 12Ψ 字节;加上 fp16 参数 2Ψ 和 fp16 梯度 2Ψ,合计 16Ψ 字节(§3.1)。GPT-2 即需 ≥24GB——是权重的 8 倍。
  2. 激活:GPT-2 在序列长 1024、batch 32 时激活约需 60GB;即便用激活检查点(约降到平方根量级、付 33% 重计算开销),也还要 ~8GB(§3.2)。而 100B 级模型即使开了检查点,batch 32 下激活仍要 ~60GB。
  3. 临时缓冲与碎片:梯度 all-reduce 前会把梯度拼成单个扁平缓冲区,1.5B 模型的 fp32 扁平缓冲就要 6GB;碎片问题更隐蔽——极端情况下显存还有超过 30% 空闲却因找不到连续块而 OOM(§3.2)。
类比:把显存想成行李箱,参数只是衣服本身,而训练还要求你为每件衣服带一个"备用衣架"(fp32副本)、一本"修改记录"(动量)和一张"波动表"(方差),且这三样必须原样复制进每个队员的箱子(DP 复制冗余)。类比失效处:行李箱装不下可以换大箱子,GPU 显存换不了;且 ZeRO 不是压缩物品,而是让每个队员只保管全部物品的 1/Nd,需要别人那份时临时借调——这是类比覆盖不到的动态调度部分。

2.2 公式手术:16Ψ 从哪来

M(模型状态) = 2Ψ(fp16参数) + 2Ψ(fp16梯度) + KΨ(优化器状态),混合精度 Adam 取 K = 12 ⇒ 共 16Ψ 字节
符号它是什么直觉
Ψ模型参数个数一切账目的基数
fp16 参数占的字节每参数 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 倍。

三、逐章精读 · ZeRO-DP:三级分区(对应论文§5、§7)

3.1 学习目标

学完本章你应能:①默写三个阶段的内存公式并解释每一项来源;②手算任意 (Ψ, Nd) 组合下的每卡模型状态内存;③说出为什么 Pos+g 不增加总通信量、而 Pp 恰好多 50%;④指出"ZeRO 改变了数学语义吗"这个问题的答案(不变——它只是重新分配了存储位置和通信时机)。

3.2 三级递进的内存公式

Pos:  4Ψ + KΨ/Nd  →  大 Nd 时 ≈ 4Ψ(4x 缩减)
Pos+g:  2Ψ + 14Ψ/Nd  →  ≈ 2Ψ(8x 缩减)
Pos+g+p:  16Ψ/Nd(随并行度线性缩减)
阶段切什么机制通信代价
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)

3.3 手算验证:Figure 1 与 Table 1 对账

取 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)的根源:加卡不仅加算力,还直接加显存。

3.4 公式手术:通信量为何是 2Ψ 和 3Ψ

基线DP:all-reduce = reduce-scatter(Ψ) + all-gather(Ψ) = 2Ψ
Pos+g:scatter-reduce(Ψ) + 步末all-gather(Ψ) = 2Ψ(恰好不变)
Pos+g+p:forward all-gather(Ψ) + backward all-gather(Ψ) + scatter-reduce(Ψ) = 3Ψ = 1.5×
符号它是什么直觉
Ψ(元素数)一次全量参数大小的数据搬运量all-gather 搬 Ψ、reduce-scatter 也搬 Ψ
Nd数据并行度分片份数;Pp 的两次 all-gather 各摊销回 Ψ

关键洞察(【解读者推断】,由 §7 推导过程支持):基线的 all-reduce 本来就等价于一次 reduce-scatter 加一次 all-gather,所以 Pos+g 只是"把本来就存在的两个操作拆开、插进不同时机",总量当然不变;Pp 多出来的是 backward 结束后为下一轮 forward 准备参数的那次额外 all-gather。

3.5 工程账单

三级分区各自缓解的是显存压力;新增的压力分别是:Pos 要求步末同步等待(延迟);Pg 引入 bucket 管理复杂度;Pp 把通信从"每步一次大批量"变成"全程细粒度流水",对调度器要求高,且小消息带宽利用率低——这就是 DeepSeek 后续工程里反复出现的债。给后文埋的债:Pos+g+p 让每卡只剩 1/Nd 的状态,于是"加卡反而更快"的超线性区出现(§10.3),但这也意味着性能对集群规模和拓扑变得敏感——这是 §18 篇 MegaScale 要偿还的债。

3.6 常见误读

  1. "ZeRO 是一种新的并行方式"——不是,它是改造过的数据并行,计算粒度和 DP 完全一样,改的只是"谁存什么、何时通信"。
  2. "Pp 的 1.5x 通信很可怕"——在大 batch 场景下通信可以被计算充分掩盖,论文实测吞吐反而更高(§10.2);代价主要出现在通信瓶颈场景。
  3. "ZeRO 能让训练变快是因为算法更好"——收敛语义完全不变(§2.3 明确说不改变优化方法);提速来自更大可用 batch 和更高算术强度,属于系统层收益。

3.7 分级自测题(本章)

  1. L1一个 30B 参数模型用混合精度 Adam 训练,模型状态总共占多少 GB?若 Nd=64 且开启 Pos,每卡占多少?
    显示答案总量 16×30=480GB。Pos:4×30 + 360/64 = 120 + 5.625 ≈ 125.6GB。
  2. L2若把 Nd 从 64 提到 512,Pos+g 下 7.5B 模型的每卡模型状态从多少变为多少?这解释了 §10.3 的什么现象?
    显示答案Nd=64:15+105/64≈16.6GB;Nd=512:15+105/512≈15.2GB。逼近 2Ψ=15GB 的下限。这说明加卡能腾出显存装更大 batch → 算术强度上升 → 吞吐超线性增长(§10.3 super-linear scalability)。
  3. L3有人说"Pos+g 既然通信量与 DP 相同,那它就是免费的"。构造一个该说法失效的场景并解释原因。
    显示答案评分要点:①总通信量相同≠时间成本相同——Pos+g 把 reduce-scatter 提前嵌入 backward(按 bucket 流水),步末还有 all-gather,若网络带宽低或 bucket 太碎,通信无法与计算重叠,尾延迟暴露;②Pos 的步末 all-gather 引入同步点,负载不均时会拖慢整体。给出任一具体失效机制并说理充分即可满分。
  4. L4仅凭"16Ψ 字节"这一个公式和"目标:训练 1T 模型",推演论文必须回答哪三个系统问题。
    显示答案评分要点:①16TB÷单卡32GB≈500 张卡起步——必须把状态切片(引出分区方案);②切片后梯度/参数如何同步——必须证明通信不爆炸(引出 §7 分析);③除模型状态外激活/缓冲/碎片仍可能爆——必须有配套的残余内存方案(引出 ZeRO-R)。三者齐备才构成完整论证链。

四、逐章精读 · ZeRO-R:残余内存三件套(对应论文§4.2、§6、§8)

4.1 Pa:激活检查点分区

失效模式先行:模型并行(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)。

4.2 CB 与 MD

CB(常量大小缓冲):把与模型大小成正比的扁平临时缓冲改成固定大小——3B 模型的 fp32 缓冲要 12GB,改成常量大缓冲后既保住带宽又不随模型膨胀(§6.2)。MD(即时碎片整理):长短生命周期张量交错是碎片根源(checkpoint 长寿命 vs 重计算激活短寿命;参数梯度长寿命 vs 激活梯度短寿命),做法是为这两类各预分配连续大块,产出即拷入(§6.3)。

4.3 一句话蒸馏与闭卷自检

一句话记住本节:ZeRO-R 管"剩下的"显存——激活靠分区(顺带救了MP的隐性复制),缓冲靠定长,碎片靠预分配。

闭卷自检清单:不看材料,我能说出——①Pa 解决的是 MP 的哪个隐性问题?②100B 模型激活检查点 33GB→2GB 依赖哪个除数?③MD 为什么把张量按生命周期分成两类?④CB 为什么不用"越大越好"的缓冲?

五、逐章精读 · 实验:ZeRO-100B(对应论文§9–§10)

配置:实现的是 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/卡)。

数字对账·一处诚实的不对称:附录表格坦承部分 baseline 实验只用了 384 或 256 卡(GPU 数须为 MP 度整数倍),且这反而给了 baseline 少卡少通信的优势;同时比较口径是"每卡吞吐"而非聚合吞吐。这是难得的自我披露,读表格时要注意 170B 行 baseline 只有 256 卡。

万亿之路的现实约束(§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 的伏笔。

六、批判性阅读:如何不被数字带节奏

6.1 证据分级

6.2 第二坐标轴

论文的叙事轴是"模型大小×吞吐",但对以下成本着墨少:Pp 未实现未测(其 1.5x 通信只是纸面分析);Pa+cpu 在多数情况下降性能、仅在极端大模型下才划算(§10.5 自己承认 C5 在 60B 上反而更慢);开发成本上 ZeRO 易用性确实高,但 MD/Pa 的调度逻辑复杂度转移到了框架内部。

6.3 论文没有告诉你什么

  1. 没有端到端收敛曲线:除 Turing-NLG 外,所有吞吐实验未报告最终模型质量,"提速"与"训得好"之间留了一道缝(尽管 §2.3 声称不改语义)。
  2. 没有与 PP 的正面对比实验:§2.1 批评 G-pipe/PipeDream 只停留在论述层面。
  3. 容错与弹性缺失:400 卡长训的故障恢复完全未提——生产级训练的真实痛点。
  4. 1T 可行性依赖 32GB V100 与理想 all-to-all 带宽:Table 1 未计入 Pp 运行时的参数临时驻留峰值与通信缓冲。

七、综合考核(毕业关)

  1. L4·重建因果链只从"16Ψ 字节"与"1T 目标"出发,重建论文的论证顺序,并说明为什么 ZeRO 先做模型状态、后做残余状态。
    显示答案与评分标准要点:①16Ψ 说明优化器状态(12Ψ)是最大单项→先切它收益最大且零通信代价(Pos);②其次梯度(2Ψ),同样零代价(Pg);③参数最后切因为要付1.5x通信(Pp);④模型状态解决后,残余状态(激活60GB级、缓冲6GB级)才成为次要瓶颈→ZeRO-R殿后。按"收益排序+代价排序"双逻辑给分。
  2. L4·数字总对账核对以下数字链是否自洽:7.5B/64卡 → 31.4/16.6/1.88GB;1T/1024卡 → 15.6GB;通信 2Ψ vs 3Ψ。
    显示答案三条链均自洽:30+90/64=31.4✓;15+105/64=16.6✓;120/64=1.875≈1.88✓;16000/1024=15.6✓;reduce-scatter+all-gather=Ψ+Ψ=2Ψ✓;再加backward前的all-gather=3Ψ=1.5×2Ψ✓。唯一口径差:16TB vs 16000GB(十进制/二进制近似)。
  3. 设计决策答辩三连问:为什么选 DP 而不是继续加深 MP?为什么 Pg 用 bucket 化?为什么不默认 offload 到 CPU?
    参考答案与评分要点①MP 切细计算粒度+跨节点带宽骤降(NVSwitch 300GB/s → IB 12.5GB/s,40B 模型实测 <5 Tflops/卡),DP 保住计算粒度;②bucket 化为重叠通信与计算、避免小消息低带宽;③PCI-E 带宽受限,offload 最高可耗 50% 训练时间(§2.2.2 引他文),只在激活检查点这类"低频访问"对象上划算。每问按"理由+反事实+论据强度"三点给分。
  4. 证据审计回看 1.3 节预评并修正。
    参考答案贡献2维持★★★(理论+实测双重);贡献3升半星——C1–C5 消融确实隔离了各组件贡献(40B→60B 归因 Pa,→140B 归因 Pos+g,→150B 归因 offload);贡献5降为半星——Pp 根本没实现,"1T"纯属内存算术;贡献6维持低分——单一困惑度指标不足以支撑 SOTA 叙事。