跳到主要内容

数据截至 (上游 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
梯度fp16
主权重(master weights)fp32
Adam 动量 mfp32
Adam 方差 vfp32
合计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 参数

三级对应 ZeroStageEnumdeepspeed/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_gradsdeepspeed/runtime/zero/stage_1_and_2.py:225zero_stage_string 在 :225)。所以本章把两者一起讲,差异处单独点出。


3. 直觉:摊平 → 等分 → 各管一段

ZeRO-1/2 的全部机关可以拆成三步直觉:

  1. 摊平(flatten):把一组参数的所有张量拼成一条连续的一维向量。参数边界消失,后面所有操作都以「元素区间」为单位。
  2. 等分(partition):把这条向量切成 Nᵈ 段,rank r 拥有第 r 段。
  3. 各管一段:每个 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_gradientsdeepspeed/runtime/zero/config.py:302)后,参数先按大小交错重排再摊平(_round_robin_reorderstage_1_and_2.py:815)。原因写在注释里(:415):反向传播按层逆序产生梯度,若连续几个梯度恰好属于同一 rank 的分片,该 rank 的 reduce 会挤在一起;交错后归主在各 rank 间轮转,通信更均。


5. 真实实现:backward 时梯度去哪了

这是 stage 2 的省钱现场,也是理解「通信量不变」的关键。

5.1 梯度桶(IPG bucket)

反向传播一开始,create_gradient_handling_hooksdeepspeed/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)把整桶一次性通信掉。桶化把几百次小张量通信合并成几次大通信,这是吞吐的来源。

5.2 stage 1 与 stage 2 的分叉点

average_tensor 里按 reduce_scatter 配置分流(:1369 起):

  • reduce_scatter 关闭:整桶 allreduce,每卡拿到全量梯度(stage 1 的典型路径)。
  • reduce_scatter 开启(默认):桶内按「这段梯度归哪个 rank」拆成 slice,每段只 reduce 到它的属主 rank——allreduce_no_retain(单属主)或 allreduce_and_scatter:1315,跨分片参数的多属主)。每卡只收 1/N 的梯度,这就是 stage 2 的梯度显存 ÷N。

通信语义对比(N 卡、梯度总量 G):

方案每卡发送+接收每卡梯度落地
DDP all-reduce2GG(全量)
ZeRO-2 reduce-scatterGG/N(只留自己段)
+ step 后 all-gather 参数G——

reduce-scatter + all-gather 合计 2G,与 all-reduce 相同——通信总量不变,显存 ÷N,这就是 ZeRO-2 的交换比。

5.3 反向结束的收尾

所有 hook 都触发后跑 independent_gradient_partition_epilogue:920):把桶里剩的尾巴 reduce 掉,然后在本 rank 分片范围内把梯度拼成一条 flat 向量 averaged_gradients[i](stage 2 走 get_flat_partition:2119)。非本分片的梯度在 reduce 时已随手释放(clear_grad_attribute:1747-1755 的 stage 分支)。


6. 真实实现:step 的全部动作

stepdeepspeed/runtime/zero/stage_1_and_2.py:2291)在梯度累积边界被引擎调用,按序做六件事:

  1. 溢出检查。 fp16 训练先 check_overflow:2520)扫全组分片梯度有没有 inf/NaN;有则 _update_scale:2523)调小 loss scale,这步直接作废(zero_grad 后 return)。
  2. 算全局梯度范数。 scaled_global_norm:2237)在各分片局部范数上 all-reduce 出全局值——注意范数也是分片算的,先平方和再开根,不还原全量梯度。
  3. 反缩放 + 裁剪。 unscale_and_clip_grads:2435)只作用于本 rank 的梯度分片。
  4. 本地 Adam。 _optimizer_step:2258)把 param_groups 临时换成只剩当前组,调真优化器的 step()。优化器状态、momentum 全是 1/N。
  5. 写回 fp16。 更新后的 fp32 分片 copy_parallel_partitioned_bit16_groups[i][partition_id]——即那条摊平向量里本 rank 负责的一段。
  6. all-gather 还原。 all_gather_dp_groupsdeepspeed/runtime/utils.py:1027)把各 rank 的新分片拼回 bit16_groups_flat,模型权重重新全员在位;_update_model_bit16_weights:796)把各参数的 .data 重新指到 flat 向量的对应视图上。

至此一个训练步完成。整步里任何时刻,单卡的优化器状态和(stage 2)梯度都只有 1/N;参数始终完整。


7. 关键细节与坑

  • 一个参数可以横跨分片。 切的是摊平后的向量,不是参数。所以代码里到处是 first_offsetgrad_start_offset 这类边界簿记(get_partition_info:1986 起)。调试看到 params_not_in_partition 不要惊讶。
  • 对齐 padding 是硬约束。 摊平长度补齐到 2 × world_size 的倍数,且每个分片起始地址要过 4 字节对齐断言(:488partitioned_data.data_ptr() % (2 * nccl_start_alignment_factor))。nccl_start_alignment_factor = 2:348)就是按 fp16 元素 4 字节对齐算出来的,改它之前先想清楚 NCCL 的对齐要求。
  • overlap_comm 的默认是动态的。 不配时由 overlap_comm_validdeepspeed/runtime/zero/config.py:383)决定——stage 3 默认开,stage 1/2 默认关(论文里 stage 1/2 的通信已经被桶化摊进 backward,显式 overlap 收益小)。
  • contiguous_gradients=False 会退回慢路。 不再用大桶,逐参数通信(buffered_reduce_fallback)。CPU offload 强制开连续梯度(:261self.contiguous_gradients = contiguous_gradients or self.cpu_offload)。
  • Muon 优化器的互斥。 reduce_scatter + optimizer offload 直接 ValueError:232);ZeRO-3 下 Muon 与 reduce_scatter 整体互斥(stage3.py:353)。用 Muon 前先读这两处断言。
  • stage 1 并不省梯度显存。 只切优化器状态;梯度仍全量 all-reduce、全量落地,只在 step 时切出本 rank 那段(:1753 的注释分支)。很多人配 stage 1 期待梯度也省,这是最常见的误解。
  • 未使用参数的处理。 ignore_unused_parameters=True(默认,zero/config.py:287)时用 hook 计数探测实际参与 backward 的参数;某些参数这步没被用到也不会挂,但跨 rank 的使用情况必须一致,否则 trace/计数错位(ZeRO-3 同理,见第 3 章)。

8. 代码地图

主题文件路径符号名
stage 枚举与配置deepspeed/runtime/zero/config.pyZeroStageEnumDeepSpeedZeroConfig.stage
stage 1/2 优化器deepspeed/runtime/zero/stage_1_and_2.pyDeepSpeedZeroOptimizer
摊平+对齐deepspeed/runtime/zero/stage_1_and_2.py:1198flatten_dense_tensors_aligned
等分切片deepspeed/runtime/zero/stage_1_and_2.py:1976get_data_parallel_partitions
换优化器视野deepspeed/runtime/zero/stage_1_and_2.py:518param_group['params'] = [single_partition_of_fp32_groups[i]]
梯度 hook 与桶deepspeed/runtime/zero/stage_1_and_2.py:1162create_gradient_handling_hooksIPGBucket
桶通信deepspeed/runtime/zero/stage_1_and_2.py:1360average_tensorallreduce_and_scatter
backward 收尾deepspeed/runtime/zero/stage_1_and_2.py:929independent_gradient_partition_epilogue
step 主流程deepspeed/runtime/zero/stage_1_and_2.py:2291step_optimizer_stepunscale_and_clip_grads
参数 all-gather 还原deepspeed/runtime/utils.py:1027all_gather_dp_groups
引擎分流入口deepspeed/runtime/engine.py:2450_configure_zero_optimizer