数据截至 (上游 commit 32e301ffaf5a)
01 · ZeRO 的显存数学与 Stage 1/2
这一章讲什么: ZeRO 到底把显存账算成了什么样,以及 Stage 1/2 在
stage_1_and_2.py里怎么实现——摊平、切片、梯度桶、本地 step、all-gather 还原。读完你能自己推导出「N 张卡能省多少显存」,也知道省出来的代价是什么。
1. 它要解决的小问题
普通数据并行(DDP)里,每张卡存一份完整的模型状态:参数、梯度、优化器状态。N 张卡就是 N 份一模一样的拷贝。
对混合精度 + Adam 训练,这份「完整拷贝」的账是(ZeRO 论文 arXiv:1910.02054 的口径,Ψ = 参数量):
| 组成部分 | 精度 | 字节数 |
|---|---|---|
| 参数(前反向用) | fp16 | 2Ψ |
| 梯度 | fp16 | 2Ψ |
| 主权重(master weights) | fp32 | 4Ψ |
| Adam 动量 m | fp32 | 4Ψ |
| Adam 方差 v | fp32 | 4Ψ |
| 合计 | 16Ψ |
13B 模型就是 208GB——单卡放不下,8 卡 DDP 每卡还是 208GB。冗余不在通信,在存储。 ZeRO 的问题因此是:能不能让每卡只持有 1/N,又不改变训练语义?
2. 三级分片:账怎么变
ZeRO 把三样状态按 DP 组大小 Nᵈ 逐层切片。每级的显存账(每卡):
| Stage | 切什么 | 每卡显存(模型状态部分) | 额外通信 |
|---|---|---|---|
| 0(关闭) | 不切 | 16Ψ | 梯度 all-reduce(同 DDP) |
| 1 | 优化器状态(fp32 主权重+m+v) | 4Ψ + 12Ψ/Nᵈ | + step 后的参数 all-gather |
| 2 | + 梯度 | 2Ψ + 14Ψ/Nᵈ | 梯度改 reduce-scatter,总量不变 |
| 3 | + 参数 | 16Ψ/Nᵈ | 前反向每步都要 all-gather 参数 |
三级对应 ZeroStageEnum(deepspeed/runtime/zero/config.py:81):disabled=0 / optimizer_states=1 / gradients=2 / weights=3。引擎按 zero_optimization.stage 分流:stage ≤2 走本章的 DeepSpeedZeroOptimizer,stage 3 走另一条实现(第 3 章),分流点在 deepspeed/runtime/engine.py:2450 _configure_zero_optimizer。
stage 1 和 2 共用同一个类,只差一个开关:partition_grads(deepspeed/runtime/zero/stage_1_and_2.py:225,zero_stage_string 在 :225)。所以本章把两者一起讲,差异处单独点出。
3. 直觉:摊平 → 等分 → 各管一段
ZeRO-1/2 的全部机关可以拆成三步直觉:
- 摊平(flatten):把一组参数的所有张量拼成一条连续的一维向量。参数边界消失,后面所有操作都以「元素区间」为单位。
- 等分(partition):把这条向量切成 Nᵈ 段,rank r 拥有第 r 段。
- 各管一段:每个 rank 只在自己那段上跑 Adam、只留自己那段的优化器状态和梯度;更新完把各段 all-gather 拼回完整参数。
通信换显存的交换比是:梯度通信总量不变(all-reduce 换成 reduce-scatter + all-gather,两次通信各半),换来优化器状 态(+梯度)显存 ÷Nᵈ。
图示:一条向量被切走(8 卡为例)
一个参数组(flatten 后的一条 fp16 向量)
┌────────────────────────────────────────────────────┐
│ seg0 │ seg1 │ seg2 │ seg3 │ seg4 │ seg5 │ seg6 │ seg7 │
└────────────────────────────────────────────────────┘
│ │ │ │ │ │ │ │
▼ ▼ ▼ ▼ ▼ ▼ ▼ ▼
rank0 rank1 rank2 rank3 rank4 rank5 rank6 rank7
各自持有: fp32 主权重分片 + Adam m/v 分片 + (stage2) 梯度分片
│
▼ step 后
all-gather:8 段新参数拼回完整向量,所有 rank 同步
原理演示
# 示意,非源码:ZeRO-1/2 的核心三步
def zero_stage2_init(params, dp_rank, dp_size):
flat = flatten(params) # ① 摊平成一条向量
segs = split_evenly(flat, dp_size) # ② 等分成 N 段
my_fp32 = segs[dp_rank].clone().float() # ③ 本地 fp32 主权重分片
optimizer = Adam([my_fp32]) # 优化器只见 1/N,状态自然 1/N
return flat, segs, my_fp32, optimizer
def zero_stage2_step(flat, segs, my_fp32, optimizer, grad_bucket):
my_grad = reduce_scatter(grad_bucket) # 只收「归我更新」的梯度段
my_fp32.grad = my_grad
optimizer.step() # 只更新 1/N
segs[dp_rank].copy_(my_fp32.half()) # 写回 fp16 分片
all_gather(flat, segs) # 凑齐完整参数供下一步用
重点看:真实优化器(Adam)从头到尾不知道 ZeRO 的存在——它只看到一个参数数为 1/N 的「小模型」。这是整个设计最省事的一刀。
4. 真实实现:初始化时发生了什么
入口是 DeepSpeedZeroOptimizer.__init__(deepspeed/runtime/zero/stage_1_and_2.py:148)。对每个优化器参数组,依次做四件事:
① 摊平。 flatten_dense_tensors_aligned(:1189)把组内参数拼成一条向量,按 alignment = 2 × world_size 补齐(4 字节对齐,给后面的 all-gather 用)。显存不够时会把参数先挪到 CPU、摊平后再搬回 GPU(:403 起的 flatten_on_accelerator 分支)。
② 等分。 get_data_parallel_partitions(:1967)用 tensor.narrow 把向量切成 dp_size 段,余数分给前几个 rank:
base_size = total_num_elements // dp
remaining = total_num_elements % dp
③ 换优化器的视野。 本 rank 的分片 clone 成 fp32 主权重 single_partition_of_fp32_groups[i],然后:
param_group['params'] = [self.single_partition_of_fp32_groups[i]]
(deepspeed/runtime/zero/stage_1_and_2.py:518)——执行完这行,真优化器的参数列表里只剩本 rank 的 1/N 分片。Adam 的动量/方差随之只在分片上分配(首次 step 时惰性初始化,:838 initialize_optimizer_states)。stage 1 的「优化器状态 ÷N」就是这一行的副作用。
④ 记录参数↔分片映射。 get_partition_info 算出哪些参数落在本分片(params_in_partition)、哪些不在(params_not_in_partition)、第一个跨界参数的偏移(first_offset)。一个参数可能横跨两个分片——切的是向量,不是参数。
4.1 round-robin 重排(可选优化)
默认按声明顺序摊平。开了 round_robin_gradients(deepspeed/runtime/zero/config.py:302)后,参数先按大小交错重排再摊平(_round_robin_reorder,stage_1_and_2.py:815)。原因写在注释里(:415):反向传播按层逆序产生梯度,若连续几个梯度恰好属于同一 rank 的分片,该 rank 的 reduce 会挤在一起;交错后归主在各 rank 间轮转,通信更均。
5. 真实实现:backward 时梯度去哪了
这是 stage 2 的省钱现场,也是理解「通信量不变」的关键。
5.1 梯度桶(IPG bucket)
反向传播一开始,create_gradient_handling_hooks(deepspeed/runtime/zero/stage_1_and_2.py:1162)给每个可训练参数注册一个 grad hook。梯度一生成,process_gradients 把它拷进一个 reduce_bucket_size(默认 5 亿元素,zero/config.py:113)大小的连续缓冲区 IPGBucket(:114)。
桶满(或 backward 结束)时,reduce_ipg_grads(:1692)→ average_tensor(:1351)把整桶一次性通信掉。桶化把几百次小张量通信合并成几次大通信,这是吞吐的来源。