L3.2 显存优化与 ZeRO
三维坐标
layer: L3(训练)|level: Senior|pillar: 训推框架上一篇我们看清了「单卡装不下大模型」这堵显存墙。本文不 再泛谈,而是把一张 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 在给定 下的单卡显存与额外通信量,解释为什么 ZeRO-1/2 是「免费午餐」、ZeRO-3 要多付 50% 通信;④ 能判断「激活」何时反超 16Ψ 成为头号 OOM,并用激活重计算把它从 降到 ;⑤ 会在 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 字节/参数从哪来
设模型参数量为 (参数个数)。在 混合精度(fp16/bf16 + Adam) 训练下,常驻显存由四部分构成。逐项记账(以 fp16 为例,每个 fp16 数 2 字节,每个 fp32 数 4 字节):
| 显存项 | 精度 | 字节/参数 | 说明 |
|---|---|---|---|
| 模型参数(fp16 副本) | fp16 | 2Ψ | 前向/反向用的工作副本 |
| 梯度(fp16) | fp16 | 2Ψ | 反向算出的梯度 |
| Optimizer:参数 master 副本 | fp32 | 4Ψ | Adam 必须用 fp32 主权重防舍入误差 |
| Optimizer:Adam 动量 m | fp32 | 4Ψ | 一阶矩 |
| Optimizer:Adam 方差 v | fp32 | 4Ψ | 二阶矩 |
把后三项(fp32 master + m + v)合起来就是 Optimizer State = 12Ψ,加上 fp16 参数 2Ψ 与梯度 2Ψ,总计:
这就是著名的 「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Ψ 按数据并行度 切成 份,每张卡只持有自己负责的那一片,需要完整数据时临时通过通信凑齐。三个阶段逐步切得更狠:
- ZeRO-1(切 Optimizer State):只把 12Ψ 的优化器状态切成 。参数与梯度仍全量冗余。
- ZeRO-2(+切 Gradient):在 ZeRO-1 基础上,把梯度也切成 。每张卡反向时只保留自己负责分片的梯度。
- ZeRO-3(+切 Parameter):连 fp16 参数本身都切成 。前向/反向用到某层权重时,临时 All-Gather 凑齐该层,算完即丢。
2.3 单卡显存与通信复杂度逐阶推导
把每张卡的常驻显存写成公式(忽略激活),设 为数据并行度:
| 阶段 | 单卡显存 | 时(7.5B,单位 GB) | 额外通信量(相对朴素 DP 的 All-Reduce 基线 ) |
|---|---|---|---|
| 朴素 DP | 120 | 1× (Reduce-Scatter + All-Gather ≈ ) | |
| ZeRO-1 | 31.4 | 1×(与 DP 同,仍是 ) | |
| ZeRO-2 | 16.6 | 1×(梯度 Reduce-Scatter + 参数 All-Gather ≈ ) | |
| ZeRO-3 | 1.9 | 1.5×(前向 All-Gather + 反向 All-Gather + 梯度 Reduce-Scatter ≈ ) |
推导要点:
- ZeRO-1:参数(2)+梯度(2)全量 = 4Ψ,优化器切片 = 。通信上,梯度仍用标准 All-Reduce(可拆为 Reduce-Scatter + All-Gather,总量 ),与朴素 DP 完全一致。这是「白嫖」——显存降 4 倍,通信零 增加。
- ZeRO-2:参数 2Ψ 全量 + 梯度/优化器切片 = 。梯度不再做全量 All-Reduce,而是 Reduce-Scatter(每卡只收自己分片的归约结果,),参数更新后 All-Gather(),总量仍 ,通信量不变。
- ZeRO-3:全部切片,单卡 ,显存随 近乎线性下降。但代价是参数被切散,前向需 All-Gather 凑参数()、反向再 All-Gather 一次()、梯度 Reduce-Scatter(),通信总 量约 ,即 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 处的激活,反向需要某段中间激活时临时重 新前向一次算出来。这是经典的 「用时间换显存」——激活显存可从 降到 ,代价是多约 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.json 的 stage 改为 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 环境下用两种方式理解原理:
- 读 ds_config 推演:把上面的
stage与 offload 配置当作「显存调度声明」逐项对照原理与架构一节的公式,手算每张卡的常驻显存——这正是 ZeRO 配置的本质。 - 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 单卡看不出省显存优势: 时 毫无切分收益,反而多付 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」,而不是无脑选显存最省的档。
用具体数字算一遍(,,1Ψ 字节 = 30 GB):
- ZeRO-1: ≈ 142.5 GB——单卡 80GB 直接 OOM,淘汰。
- ZeRO-2: ≈ 86.25 GB——略超 80GB,这 86.25 GB 是与激活无关的常驻状态,单靠
overlap_comm压不掉;要用 ZeRO-2 就必须叠加 CPU offload 或缩小常驻(如优化器状态 offload),工程上属于「压线可行但要打补丁」。 - ZeRO-3: ≈ 30 GB——最省,单卡余量充足。
怎么选:显存目标是「放得下 + 留足激活余量」,不是「越小越好」。ZeRO-3 虽把常驻状态压到 30GB,但它的代价写在通信账上——前向 All-Gather + 反向 All-Gather + 梯度 Reduce-Scatter ,总量约 ,是朴素 DP()的 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 的 通信复杂度,解释瓶颈在哪一步通信、为什么是跨节点链路,以及你会优先尝试哪些手段(如 overlap_comm、退回 ZeRO-2 + 激活重计算)?
展开参考答案(含跨节点通信瓶颈链路图 + 带宽算一遍)
结论:ZeRO-3 每步要做三次集合通信(前向 All-Gather、反向 All-Gather、梯度 Reduce-Scatter),共约 流量;这些流量一旦跨节点,就被压到 100Gb 以太网这条比 NVLink 慢一两个数量级的细管子里,链路带宽成为木桶最短板,于是计算等通信、吞吐腰斩。
用带宽算一遍(数量级估算):
- NVLink(如 A100 第三代)单卡聚合带宽约 600 GB/s 级;100Gb 以太网理想吞吐约
100/8 = 12.5 GB/s,实际打折后更低——两者差约 40~50 倍。 - ZeRO-3 每步通信量约 ;设某层参数分片需在节点间 All-Gather 的流量为 字节,节点内只需 秒,跨节点却要 秒——同一笔流量,跨节点耗时是节点内的几十倍。
- 后果:当 All-Gather 必须等最慢的跨节点链路返回参数才能开始算这一层,计算单元被迫空转等数据,
overlap_comm也只能重叠掉一部分——带宽不够时,通信时间长到算力根本掩盖不住,吞吐随之腰斩。
优先尝试的手段(从低成本到高成本):
| 手段 | 在做什么 | 适用前提 |
|---|---|---|
开 overlap_comm: true | 让 All-Gather / Reduce-Scatter 与计算重叠,掩盖部分延迟 | 通信量未远超算力时有效 |
| 退回 ZeRO-2 + 激活重计算 | 把通信从 降回 ,省下的显存靠重计算补回 | 显存能放下 ZeRO-2 时首选 |
| 调整并行拓扑 | 把 ZeRO 通信尽量限制在节点内 NVLink、跨节点改走 TP/PP | 需要 3D 并行编排(见思考题 3) |
| 升级互联 | 100Gb 以太网换 InfiniBand / RoCE | 有硬件预算时的根治手段 |
为什么是跨节点链路:瓶颈不在通信「次数」而在「最慢的那一跳」。ZeRO-3 比 ZeRO-2 多出的那份 All-Gather 通信量( 与 之差),在 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])按列/行切成 片,每张 TP 卡常驻自己那一片,前向用 All-Reduce 拼接激活——参数已经是「切开且分布式持有」的状态。 - ZeRO-3 的做法:沿数据并行维度把参数再切成 片,用到某层时临时 All-Gather 把整层参数凑齐、算完即丢。
- 冲突点:ZeRO-3 的 All-Gather 假设「全量参数本应在某处可凑齐」,但 TP 下整层参数本就被设计成永不在单卡凑齐(凑齐就违背了 TP 省显存的初衷)。两者对同一份权重的「切」与「凑」语义相互打架——要么 ZeRO-3 把 TP 已切的再切一遍造成元数据与通信冗余,要么需要极复杂的嵌套编排才能自洽,收益却很小。
为什么 TP/PP + ZeRO-1 是主流配方:
| 组合 | 切谁 | 维度关系 | 工程结果 |
|---|---|---|---|
| TP/PP + ZeRO-1 | TP/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 多卡 DP | ZeRO-1 或 2 | stage 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 或吞吐低)
- OOM:降
micro_batch→ 升 stage(1→2→3)→ 开activation_checkpointing→ offload - 吞吐低:开
overlap_comm→ 调大reduce_bucket_size/allgather_bucket_size→ 检查是否跨节点 ZeRO-3 - loss NaN:查 bf16 是否需
loss_scale(fp16)或 grad clip - resume:DeepSpeed checkpoint 含 optimizer sharding 状态,勿只存
model.pt
4.4 与 FSDP 对照
| 项 | DeepSpeed ZeRO | PyTorch FSDP |
|---|---|---|
| 配置 | ds_config.json | FullyShardedDataParallel(...) |
| Stage 3 等价 | "stage": 3 | ShardingStrategy.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.py与stage3.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 并行、数据流水线协同落地。