跳到主要内容

L3.2 显存优化与 ZeRO

三维坐标 layer: L3(训练)level: Seniorpillar: 训推框架

上一篇我们看清了「单卡装不下大模型」这堵显存墙。本文不再泛谈,而是把一张 GPU 上的每一字节显存都记成账——参数、梯度、优化器状态、激活各占多少;再用 ZeRO(Zero Redundancy Optimizer) 把冗余的副本沿数据并行维度切开,让显存随 GPU 数近乎线性下降。读完你应能徒手推导「7B 模型为什么 16 字节/参数」、「ZeRO-3 多付了多少通信量」,并在 DeepSpeed 里一行配置切换三档验证。

学习目标

  • 前置知识:读过 L0.3 显存账本(知道一个浮点数占几个字节、fp16/fp32 的差异)与 L3.1 并行(理解 DP / TP / PP 三种切分维度的区别);会读简单的数学公式与 JSON 配置;知道「数据并行就是每张卡各跑一份模型、再同步梯度」即可。无需读过 DeepSpeed 源码。
  • 学完产出:① 能徒手推导「fp16 + Adam = 16Ψ 字节/参数」这本显存账本,说清参数 / 梯度 / 优化器状态各占几 Ψ、为什么 master 权重必须 fp32;② 能画出 ZeRO-1/2/3 沿 DP 维度逐阶切 optimizer state → gradient → parameter 的切分图,并写出每张卡的单卡显存公式;③ 能算出三档 stage 在给定 NdN_d 下的单卡显存与额外通信量,解释为什么 ZeRO-1/2 是「免费午餐」、ZeRO-3 要多付 50% 通信;④ 能判断「激活」何时反超 16Ψ 成为头号 OOM,并用激活重计算把它从 O(L)O(L) 降到 O(L)O(\sqrt{L});⑤ 会在 DeepSpeed 里只改 ds_config.json 一个字段切换三档 stage,实测峰值显存与吞吐验证理论。
  • 阅读姿势:盯住一条主线——「ZeRO 的一切都是在用通信换显存:把数据并行里每张卡冗余存的那份『参数 + 梯度 + 优化器状态』沿 DP 维度切开,要用时再临时通信凑齐」。从 16Ψ 这本账本出发,每多切一项,就多省一截显存、多付一点通信,stage 选型本质就是在这条「显存 vs 通信」的天平上找平衡点。

背景与现状

大模型训练的第一性矛盾,从来不是算力不够,而是显存放不下。一块 80GB 的 A100/H100,朴素地用 Adam + fp16 训练,能塞下的稠密模型规模只有约 3~4B 参数——这与动辄百亿、千亿的模型规模差了一两个数量级。这道墙的学名叫 显存墙(Memory Wall)

业界破墙的思路分两条路线:

  • 「切模型」路线:张量并行(TP)、流水线并行(PP)把单个模型的权重/计算切到多卡。代价是改动模型代码、引入复杂的通信编排,对算子有侵入性。
  • 「切状态」路线:保持数据并行(DP)的简单编程模型不变,但发现 DP 的本质是每张卡冗余地存了一份完整的「参数 + 梯度 + 优化器状态」。既然冗余,就能切——这正是 ZeRO 的核心 insight:用通信换显存,且不改一行模型代码

业界信号:ZeRO 由微软于 2020 年在论文 ZeRO: Memory Optimizations Toward Training Trillion Parameter Models 提出,落地为 microsoft/DeepSpeed,并被 PyTorch 官方吸收为 FSDP(Fully Sharded Data Parallel)——二者思想同源。今天几乎所有大模型团队的训练底座,要么是 DeepSpeed/ZeRO,要么是 FSDP,要么是 Megatron + ZeRO 的组合(3D 并行)。ZeRO 能成为「显存优化事实标准」,正是因为它几乎零侵入:把单卡 DP 脚本套上 ZeRO,无需重写模型即可线性扩容。

理解 ZeRO 的前提,是先看懂混合精度训练这本显存账本——这是原理与架构一节的起点。

原理与架构

2.1 混合精度训练的显存账本:16 字节/参数从哪来

设模型参数量为 Ψ\Psi(参数个数)。在 混合精度(fp16/bf16 + Adam) 训练下,常驻显存由四部分构成。逐项记账(以 fp16 为例,每个 fp16 数 2 字节,每个 fp32 数 4 字节):

显存项精度字节/参数说明
模型参数(fp16 副本)fp16前向/反向用的工作副本
梯度(fp16)fp16反向算出的梯度
Optimizer:参数 master 副本fp32Adam 必须用 fp32 主权重防舍入误差
Optimizer:Adam 动量 mfp32一阶矩
Optimizer:Adam 方差 vfp32二阶矩

把后三项(fp32 master + m + v)合起来就是 Optimizer State = 12Ψ,加上 fp16 参数 2Ψ 与梯度 2Ψ,总计:

Mmodel+grad+optim=2Ψ+2Ψ+12Ψ=16Ψ 字节M_{\text{model+grad+optim}} = 2\Psi + 2\Psi + 12\Psi = 16\Psi \ \text{字节}

这就是著名的 「fp16 + Adam = 16 字节/参数」。一个 7.5B 模型仅这部分就需 7.5e9 × 16 ≈ 120 GB——单张 80GB 卡直接 OOM。注意这还不含激活(Activation),激活随 batch size 与序列长度增长,往往是另一个显存大头(见 2.4)。

值得注意的是:在 DP=N 的朴素数据并行里,上面这 16Ψ 在每张卡上各存一份完整副本——N 张卡存了 N 份一模一样的优化器状态。这就是 「Zero Redundancy」要消灭的冗余

2.2 ZeRO 三阶切分:把 16Ψ 沿 DP 维度切开

ZeRO 的做法是把这 16Ψ 按数据并行度 NdN_d 切成 NdN_d 份,每张卡只持有自己负责的那一片,需要完整数据时临时通过通信凑齐。三个阶段逐步切得更狠:

  • ZeRO-1(切 Optimizer State):只把 12Ψ 的优化器状态切成 12Ψ/Nd12\Psi/N_d。参数与梯度仍全量冗余。
  • ZeRO-2(+切 Gradient):在 ZeRO-1 基础上,把梯度也切成 2Ψ/Nd2\Psi/N_d。每张卡反向时只保留自己负责分片的梯度。
  • ZeRO-3(+切 Parameter):连 fp16 参数本身都切成 2Ψ/Nd2\Psi/N_d。前向/反向用到某层权重时,临时 All-Gather 凑齐该层,算完即丢

2.3 单卡显存与通信复杂度逐阶推导

把每张卡的常驻显存写成公式(忽略激活),设 NdN_d 为数据并行度:

阶段单卡显存Nd=64N_d{=}64 时(7.5B,单位 GB)额外通信量(相对朴素 DP 的 All-Reduce 基线 2Ψ2\Psi
朴素 DP16Ψ16\Psi1201× (Reduce-Scatter + All-Gather ≈ 2Ψ2\Psi
ZeRO-14Ψ+12ΨNd4\Psi + \dfrac{12\Psi}{N_d}31.4(与 DP 同,仍是 2Ψ2\Psi
ZeRO-22Ψ+14ΨNd2\Psi + \dfrac{14\Psi}{N_d}16.6(梯度 Reduce-Scatter + 参数 All-Gather ≈ 2Ψ2\Psi
ZeRO-316ΨNd\dfrac{16\Psi}{N_d}1.91.5×(前向 All-Gather Ψ\Psi + 反向 All-Gather Ψ\Psi + 梯度 Reduce-Scatter Ψ\Psi3Ψ3\Psi

推导要点

  1. ZeRO-1:参数(2)+梯度(2)全量 = 4Ψ,优化器切片 = 12Ψ/Nd12\Psi/N_d。通信上,梯度仍用标准 All-Reduce(可拆为 Reduce-Scatter + All-Gather,总量 2Ψ2\Psi),与朴素 DP 完全一致。这是「白嫖」——显存降 4 倍,通信零增加。
  2. ZeRO-2:参数 2Ψ 全量 + 梯度/优化器切片 = 14Ψ/Nd14\Psi/N_d。梯度不再做全量 All-Reduce,而是 Reduce-Scatter(每卡只收自己分片的归约结果,Ψ\Psi,参数更新后 All-Gather(Ψ\Psi),总量仍 2Ψ2\Psi通信量不变
  3. ZeRO-3:全部切片,单卡 16Ψ/Nd16\Psi/N_d显存随 NdN_d 近乎线性下降。但代价是参数被切散,前向需 All-Gather 凑参数(Ψ\Psi)、反向再 All-Gather 一次(Ψ\Psi)、梯度 Reduce-Scatter(Ψ\Psi,通信总量约 3Ψ3\Psi,即 1.5 倍于朴素 DP。

要点在于:ZeRO-1/2 是几乎免费的午餐——通信量与朴素 DP 持平却大幅省显存,应作为默认起点。ZeRO-3 才是「用通信换显存」的真正分水岭:它能把万亿参数切到可训,但多付 50% 通信量,在带宽不足(如跨节点 PCIe/以太网而非 NVLink/IB)时吞吐会明显下滑。这正是 stage 选型的核心权衡:显存够用就别上 ZeRO-3。

2.4 激活:另一个被忽视的显存大户,用「重计算」换它

上面 16Ψ 全是与 batch 无关的常驻显存。但训练时还有 激活(Activation)——前向产生、反向要用的中间张量,其显存正比于 batch × seq_len × hidden × layers,在长序列大 batch 下常常反超 16Ψ 成为头号 OOM 元凶。

破解手段是 激活重计算 / 梯度检查点(Gradient Checkpointing / Activation Recomputation):前向时只保存少量 checkpoint 处的激活,反向需要某段中间激活时临时重新前向一次算出来。这是经典的 「用时间换显存」——激活显存可从 O(L)O(L) 降到 O(L)O(\sqrt{L}),代价是多约 33% 的前向计算量。ZeRO 切常驻状态,重计算切激活,二者正交,生产配置里通常同时开启

动手实践:DeepSpeed 切 ZeRO stage 实测

实验目标:在 DeepSpeed 中仅修改 ds_config.json 里的一个字段 zero_optimization.stage(1/2/3),训练同一个小模型,记录 torch.cuda.max_memory_allocated()(峰值显存)与吞吐(samples/s),亲手验证模块 2 的「显存随 stage 递减、ZeRO-3 吞吐因通信下滑」的结论。产出物:一张 stage-vs-显存-vs-吞吐 对照表。

3.1 环境准备

# 推荐 Python 3.11,用 venv/uv 隔离
python3 -m venv .venv && source .venv/bin/activate

# NVIDIA GPU 路径(推荐,能真实测显存)
pip install torch --index-url https://download.pytorch.org/whl/cu121
pip install deepspeed

# 验证 DeepSpeed 环境(会打印 CUDA / NCCL / op builder 状态)
ds_report

3.2 配置:ds_config.json(切 stage 只改一个数)

{
"train_batch_size": 4,
"train_micro_batch_size_per_gpu": 4,
"gradient_accumulation_steps": 1,
"fp16": { "enabled": true },
"zero_optimization": {
"stage": 2,
"allgather_partitions": true,
"reduce_scatter": true,
"overlap_comm": true,
"contiguous_gradients": true
},
"optimizer": {
"type": "Adam",
"params": { "lr": 1e-4 }
}
}

实验时只把 "stage" 改成 1 / 2 / 3 三次,其余不动。若要在单卡上验证 ZeRO-3 的 CPU offload 原理,可在 zero_optimization 里追加 "offload_optimizer": {"device": "cpu"}"offload_param": {"device": "cpu"}

3.3 代码:训练小模型并记录峰值显存与吞吐

import time, json, torch, torch.nn as nn, deepspeed

# 一个故意「优化器状态偏重」的小模型,便于观测 ZeRO 切分效果
class Net(nn.Module):
def __init__(self, d=4096, layers=8):
super().__init__()
self.net = nn.Sequential(*[nn.Linear(d, d) for _ in range(layers)])
def forward(self, x):
return self.net(x)

def main():
model = Net()
engine, _, _, _ = deepspeed.initialize(
model=model,
model_parameters=model.parameters(),
config="ds_config.json",
)
dev = engine.device
torch.cuda.reset_peak_memory_stats(dev)

steps, bs, d = 20, 4, 4096
loss_fn = nn.MSELoss()
t0 = time.time()
for _ in range(steps):
x = torch.randn(bs, d, device=dev, dtype=torch.float16)
y = torch.randn(bs, d, device=dev, dtype=torch.float16)
loss = loss_fn(engine(x), y)
engine.backward(loss) # ZeRO 在此做 Reduce-Scatter
engine.step() # 在此做参数更新 + All-Gather
torch.cuda.synchronize(dev)
dt = time.time() - t0

peak_gb = torch.cuda.max_memory_allocated(dev) / 1024**3
thpt = steps * bs / dt
stage = json.load(open("ds_config.json"))["zero_optimization"]["stage"]
print(f"[ZeRO-{stage}] peak_mem={peak_gb:.2f} GB throughput={thpt:.1f} samples/s")

if __name__ == "__main__":
main()

3.4 运行与观察

# 单机多卡(如 2 张 GPU),DeepSpeed 启动器自动起多进程
deepspeed --num_gpus=2 train_zero.py

依次把 ds_config.jsonstage 改为 1→2→3 各跑一次,你应观测到类似规律(绝对值随硬件而异,重点看趋势):

stage峰值显存(相对)吞吐(samples/s)解读
ZeRO-1最高最高只切优化器,通信≈DP,最快
ZeRO-2≈ ZeRO-1加切梯度,通信仍≈DP,吞吐基本持平
ZeRO-3最低略降加切参数,多 All-Gather 通信,吞吐下滑

结论与原理推导一致:显存 ZeRO-3 < ZeRO-2 < ZeRO-1;吞吐 ZeRO-3 因 1.5× 通信而最低,ZeRO-1/2 基本持平。

3.5 Mac / 无 GPU 替代方案

DeepSpeed 的 ZeRO 依赖 CUDA + NCCL,Mac/纯 CPU 无法真实测 GPU 显存。在无 GPU 环境下用两种方式理解原理:

  1. 读 ds_config 推演:把上面的 stage 与 offload 配置当作「显存调度声明」逐项对照原理与架构一节的公式,手算每张卡的常驻显存——这正是 ZeRO 配置的本质。
  2. CPU offload 演示原理:在 zero_optimization 里加 "offload_optimizer": {"device": "cpu"},DeepSpeed 会把 12Ψ 优化器状态搬到主机内存,GPU 只留参数+梯度。这与 ZeRO 切分思想一致——「不常用的状态搬到更便宜、更大的存储层级」。即便单卡也能跑通流程,理解「显存→内存→NVMe」的分层 offload 链条(ZeRO-Infinity)。

踩坑预警 (Gotchas)

  • max_memory_allocated 只统计 PyTorch 分配器:NCCL 通信缓冲、CUDA context 等不计入,真实显存请同时看 nvidia-smi。测前务必 reset_peak_memory_stats()
  • ZeRO-3 单卡看不出省显存优势Nd=1N_d{=}116Ψ/1=16Ψ16\Psi/1 = 16\Psi 毫无切分收益,反而多付 All-Gather 开销。ZeRO-3 必须多卡才有意义
  • 小模型 ZeRO-3 可能更慢且更费:参数太小时,All-Gather 通信开销 > 省下的显存收益,且分片元数据有固定开销。ZeRO-3 是给「单卡装不下」的大模型用的,别拿小模型证伪它。
  • 跨节点带宽是 ZeRO-3 杀手:NVLink/InfiniBand 上 ZeRO-3 尚可,一旦掉到 PCIe/以太网,1.5× 通信量会让吞吐断崖式下跌。开 overlap_comm: true 让通信与计算重叠以缓解。
  • train_batch_size 必须自洽:DeepSpeed 要求 train_batch_size = micro_batch × grad_accum × num_gpus,对不上会直接报错。

深入思考

下面三题每题先给题干,再用 <details> 折叠一份图文并茂的参考答案。建议先合上答案自己想 3 分钟,再展开对照。

思考题 1:stage 选型决策

你有 16 张 A100-80GB(NVLink 互联),要训练一个 30B 稠密模型。请用 2.3 节的单卡显存公式估算 ZeRO-1/2/3 各自的单卡常驻显存(不含激活),判断哪一档「显存够用且吞吐最优」,并说明为什么不应无脑选显存最省的 ZeRO-3。

展开参考答案(含 stage 选型决策图 + 算一遍)

结论:先按显存公式算出每档单卡占用,找到「第一个能放下」的 stage 就停手——本例 ZeRO-2 算出来约 86 GB、略超 80GB 单卡,需配合激活重计算/offload 等手段勉强压入;若不想引入这些补丁,严格按公式应选 ZeRO-3(约 30 GB、余量充足),代价是多付 50% 通信。选型核心是「第一个能放下的 stage」,而不是无脑选显存最省的档。

用具体数字算一遍Ψ=30e9\Psi = 30\text{e}9Nd=16N_d = 16,1Ψ 字节 = 30 GB):

  1. ZeRO-14Ψ+12ΨNd=4Ψ+0.75Ψ=4.75Ψ4\Psi + \dfrac{12\Psi}{N_d} = 4\Psi + 0.75\Psi = 4.75\Psi142.5 GB——单卡 80GB 直接 OOM,淘汰。
  2. ZeRO-22Ψ+14ΨNd=2Ψ+0.875Ψ=2.875Ψ2\Psi + \dfrac{14\Psi}{N_d} = 2\Psi + 0.875\Psi = 2.875\Psi86.25 GB——略超 80GB,这 86.25 GB 是与激活无关的常驻状态,单靠 overlap_comm 压不掉;要用 ZeRO-2 就必须叠加 CPU offload 或缩小常驻(如优化器状态 offload),工程上属于「压线可行但要打补丁」。
  3. ZeRO-316ΨNd=Ψ\dfrac{16\Psi}{N_d} = \Psi30 GB——最省,单卡余量充足。

怎么选:显存目标是「放得下 + 留足激活余量」,不是「越小越好」。ZeRO-3 虽把常驻状态压到 30GB,但它的代价写在通信账上——前向 All-Gather Ψ\Psi + 反向 All-Gather Ψ\Psi + 梯度 Reduce-Scatter Ψ\Psi,总量约 3Ψ3\Psi,是朴素 DP(2Ψ2\Psi)的 1.5 倍。本例 16 卡虽是 NVLink,通信仍非零成本。因此:ZeRO-3 单卡显存小于 ZeRO-2,但吞吐也更低(此处「小于」即 ZeRO-3 显存 < ZeRO-2 显存)——本例的取舍是:若愿意为 ZeRO-2 叠加 offload/重计算补丁把 86 GB 压进 80GB,就能省下那 50% 通信换吞吐;若追求配置简单稳妥,ZeRO-3 是「一定放得下」的兜底档。「显存够用就别上更高的 stage」,这正是 2.3 节强调的选型核心。

思考题 2:通信瓶颈定位

同样的 ZeRO-3 配置,在单机 8 卡 NVLink 上吞吐良好,迁到「2 节点 × 4 卡、节点间走 100Gb 以太网」后吞吐腰斩。结合 2.3 节 ZeRO-3 的 3Ψ3\Psi 通信复杂度,解释瓶颈在哪一步通信、为什么是跨节点链路,以及你会优先尝试哪些手段(如 overlap_comm、退回 ZeRO-2 + 激活重计算)?

展开参考答案(含跨节点通信瓶颈链路图 + 带宽算一遍)

结论:ZeRO-3 每步要做三次集合通信(前向 All-Gather、反向 All-Gather、梯度 Reduce-Scatter),共约 3Ψ3\Psi 流量;这些流量一旦跨节点,就被压到 100Gb 以太网这条比 NVLink 慢一两个数量级的细管子里,链路带宽成为木桶最短板,于是计算等通信、吞吐腰斩。

用带宽算一遍(数量级估算):

  1. NVLink(如 A100 第三代)单卡聚合带宽约 600 GB/s 级;100Gb 以太网理想吞吐约 100/8 = 12.5 GB/s,实际打折后更低——两者差约 40~50 倍
  2. ZeRO-3 每步通信量约 3Ψ3\Psi;设某层参数分片需在节点间 All-Gather 的流量为 XX 字节,节点内只需 X/600X/600 秒,跨节点却要 X/12.5X/12.5 秒——同一笔流量,跨节点耗时是节点内的几十倍
  3. 后果:当 All-Gather 必须等最慢的跨节点链路返回参数才能开始算这一层,计算单元被迫空转等数据,overlap_comm 也只能重叠掉一部分——带宽不够时,通信时间长到算力根本掩盖不住,吞吐随之腰斩。

优先尝试的手段(从低成本到高成本):

手段在做什么适用前提
overlap_comm: true让 All-Gather / Reduce-Scatter 与计算重叠,掩盖部分延迟通信量未远超算力时有效
退回 ZeRO-2 + 激活重计算把通信从 3Ψ3\Psi 降回 2Ψ2\Psi,省下的显存靠重计算补回显存能放下 ZeRO-2 时首选
调整并行拓扑把 ZeRO 通信尽量限制在节点内 NVLink、跨节点改走 TP/PP需要 3D 并行编排(见思考题 3)
升级互联100Gb 以太网换 InfiniBand / RoCE有硬件预算时的根治手段

为什么是跨节点链路:瓶颈不在通信「次数」而在「最慢的那一跳」。ZeRO-3 比 ZeRO-2 多出的那份 All-Gather 通信量(3Ψ3\Psi2Ψ2\Psi 之差),在 NVLink 上可忽略,在以太网上却被放大成主导项——这正是 2.3 节「跨节点带宽是 ZeRO-3 杀手」的实测体现:ZeRO-3 通信量小于带宽承载力时无碍,一旦通信需求 > 链路供给,吞吐就崩。

思考题 3:ZeRO 与 TP/PP 的正交组合

ZeRO 沿 数据并行维度 切状态,张量并行(TP)沿 模型维度 切权重。当显存仍不够(如训万亿模型)时,工程上常用 ZeRO + TP + PP 的 3D 并行。请论证:为什么 ZeRO-3 与 TP 在同一份权重上会产生切分冲突/冗余,业界为何更常见「TP/PP 切模型 + ZeRO-1 切优化器」而非「TP + ZeRO-3」的组合?

展开参考答案(含 3D 并行分工图 + 对比表)

结论:TP 已经把每一层的权重沿模型维度切碎并分给不同卡,ZeRO-3 又要沿数据维度对「同一份参数」再切一次并临时 All-Gather 凑齐——两者抢着切同一块权重,通信编排互相打架;而 ZeRO-1 只切优化器状态、不碰参数本身,与 TP 的模型切分天然正交,所以「TP/PP 切模型 + ZeRO-1 切优化器」成了业界默认的 3D 并行配方。

为什么 TP + ZeRO-3 会冲突/冗余

  • TP 的做法:把一层的权重矩阵(如 [hidden, 4·hidden])按列/行切成 NtN_t 片,每张 TP 卡常驻自己那一片,前向用 All-Reduce 拼接激活——参数已经是「切开且分布式持有」的状态。
  • ZeRO-3 的做法:沿数据并行维度把参数再切成 NdN_d 片,用到某层时临时 All-Gather 把整层参数凑齐、算完即丢
  • 冲突点:ZeRO-3 的 All-Gather 假设「全量参数本应在某处可凑齐」,但 TP 下整层参数本就被设计成永不在单卡凑齐(凑齐就违背了 TP 省显存的初衷)。两者对同一份权重的「切」与「凑」语义相互打架——要么 ZeRO-3 把 TP 已切的再切一遍造成元数据与通信冗余,要么需要极复杂的嵌套编排才能自洽,收益却很小。

为什么 TP/PP + ZeRO-1 是主流配方

组合切谁维度关系工程结果
TP/PP + ZeRO-1TP/PP 切参数,ZeRO-1 只切优化器状态正交:一个管权重、一个管优化器✅ 干净解耦,主流默认
TP + ZeRO-3两者都想切参数同维冲突:抢切同一份权重⚠️ 编排复杂、通信冗余、收益小

核心逻辑:3D 并行的精髓是让每个维度各切一类东西、互不重叠——PP 切层、TP 切层内权重、ZeRO 切数据冗余。ZeRO-1 只动「优化器状态」这个 TP/PP 都没碰的 12Ψ,于是能与模型切分完美叠加:TP/PP 先把万亿参数本体切到可放下,ZeRO-1 再把每个 DP 副本里冗余的优化器状态切掉,二者增益相乘而非相消。而 ZeRO-3 与 TP 都盯着「参数」这同一刀,自然不如「各切各的」来得清爽——这就是业界更偏好「TP/PP 切模型 + ZeRO-1 切优化器」的根本原因(详见 L3.1 并行与 L3.3 框架实战)。

DeepSpeed 配置调优 Playbook

本节把 ZeRO 理论落成 可复制的 ds_config.json 决策树——改一个字段前知道会影响显存、通信还是吞吐。

4.1 Stage 选型速查

场景推荐 stage关键配置
7B 单卡微调ZeRO-2 + offload 可选"stage": 2, "offload_optimizer": {"device": "cpu"}
7B 多卡 DPZeRO-1 或 2stage 1 通信最少;显存紧升 stage 2
70B 多卡ZeRO-3 或 FSDP"stage": 3, "overlap_comm": true
70B + 跨节点以太网ZeRO-2 + TP 节点内避免 ZeRO-3 跨节点 All-Gather(见思考题 2)
极限显存ZeRO-3 + offload"offload_param"/"offload_optimizer"

4.2 最小可运行配置模板

{
"train_batch_size": 32,
"train_micro_batch_size_per_gpu": 4,
"gradient_accumulation_steps": 2,
"bf16": {"enabled": true},
"zero_optimization": {
"stage": 2,
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 500000000
},
"gradient_clipping": 1.0,
"steps_per_print": 10,
"wall_clock_breakdown": true
}

启动:

deepspeed --num_gpus=2 train.py --deepspeed ds_config.json

4.3 调参顺序(遇 OOM 或吞吐低)

  1. OOM:降 micro_batch → 升 stage(1→2→3)→ 开 activation_checkpointing → offload
  2. 吞吐低:开 overlap_comm → 调大 reduce_bucket_size / allgather_bucket_size → 检查是否跨节点 ZeRO-3
  3. loss NaN:查 bf16 是否需 loss_scale(fp16)或 grad clip
  4. resume:DeepSpeed checkpoint 含 optimizer sharding 状态,勿只存 model.pt

4.4 与 FSDP 对照

DeepSpeed ZeROPyTorch FSDP
配置ds_config.jsonFullyShardedDataParallel(...)
Stage 3 等价"stage": 3ShardingStrategy.FULL_SHARD
生态Megatron 集成深torch native,2.x 默认推荐

详细框架集成见 L3.3 主流框架实践

延伸阅读

1. 核心 Paper

  • ZeRO: Memory Optimizations Toward Training Trillion Parameter Models(Rajbhandari et al., 2020)— 本篇理论基石,ZeRO-1/2/3 显存与通信分析的原始出处,必读
  • ZeRO-Offload: Democratizing Billion-Scale Model Training(2021)— 把优化器状态 offload 到 CPU,单卡也能训大模型的原理。
  • ZeRO-Infinity: Breaking the GPU Memory Wall for Extreme Scale Deep Learning(2021)— 显存→内存→NVMe 三级 offload,理解 3.5 节 offload 链条的终极形态。

2. 相关高 Star 仓库与源码必读路径

  • microsoft/DeepSpeed — 看 deepspeed/runtime/zero/stage_1_and_2.pystage3.py,对照本文公式理解切分与通信的代码实现。
  • pytorch/pytorch(FSDP)— torch/distributed/fsdp/,理解 ZeRO 思想在 PyTorch 原生侧的等价实现(FullShard ≈ ZeRO-3)。

3. 优质博客 / 视频

  • Microsoft Research 官方博客「ZeRO & DeepSpeed」系列,含直观的切分动画。
  • HuggingFace 文档「Model training anatomy」与「DeepSpeed Integration」,把 16 字节/参数的账本讲得极清楚。

下一篇L3.3 主流框架与工具实践:把 ZeRO 放回 DeepSpeed / Megatron / FSDP 的完整工程框架里,看显存优化如何与 3D 并行、数据流水线协同落地。