数据截至 (上游 commit 32e301ffaf5a)
03 · ZeRO-3:参数分片与 all-gather 调度
这一章讲什么: Stage 3 是 ZeRO 的最后一块拼图——连参数本身也按 rank 切片,每卡只剩 16Ψ/N。实现它和 stage 1/2 是完全不同的技术路线:不动优化器的视野,而是让参数在「整块」与「分片」两种形态间反复横跳,并用 hook + 预取把通信藏进计算里。
1. 它要解决 的小问题
Stage 2 之后,每卡仍持有完整参数(2Ψ 字节)+ 完整的模型构造过程。两个后果:
- 参数本身就能撑爆显存。 70B 的 fp16 参数是 140GB,单卡放不下——根本轮不到谈优化器状态。
- 模型都构造不出来。 单机 8 进程各自实例化一个 1T 模型,CPU 内存先要 8 份全量(ZeRO 官方 docstring 的算法:
partition_parameters.py:940起)。
ZeRO-3 的回答:参数也切成 N 份,每张卡平时只握着自己的 1/N 分片;前向/反向跑到某个模块时,把它的参数临时 all-gather 成整块,算完立刻切回去。显存降到 16Ψ/N,代价是参数通信量翻倍(前向、反向各 all-gather 一遍),以及一个全新的调度问题:
怎么在「刚好要用之前」把参数取回来——太早浪费显存,太晚卡住计算?
2. 直觉:参数的两种形态 + 一本「取用剧本」
每个参数在 ZeRO-3 里有两个形态,由 ZeroParamStatus(deepspeed/runtime/zero/partition_parameters.py:236)标记:
| 状态 | 含义 |
|---|---|
AVAILABLE | 完整参数在卡上,可以参与计算 |
NOT_AVAILABLE | 只有 1/N 分片(在 param.ds_tensor 里),param.data 无意义 |
INFLIGHT | all-gather 已发出、尚未完成 |
调度问题的解法很「工程」:不做静态图分析,而是跑一次记下来。 第一遍 forward/backward 时记录模块执行顺序(trace);从第二遍起,照着这本「剧本」提前预取后面的参数、推迟释放还会复用的参数。
图示:一个模块的前向(ZeRO-3)
时间 ──────────────────────────────────────────────►
模块 i-1 收尾 模块 i 模块 i+1
│ ┌──────────────────┐ │
│ 预取 i+1 │ ① 等 i 的参数到位 │ │
│ (异步) │ ② 正常算 forward │ 预取 i+2
▼ │ ③ 算完释放 i 参数 │ (异步)
release i-1 └──────────────────┘
▲
all-gather i 的参数(已在 ② 前异步发出)
原理演示
# 示意,非源码:ZeRO-3 的 fetch-compute-release 循环
def forward_with_zero3(module, x):
for layer in module.layers: # 按 trace 剧本走
all_gather_async(next_layer.params) # 预取:提前发通信
wait(all_gather(layer.params)) # ① 当前层参数到位
x = layer(x) # ② 正常计算
if not reused_soon(layer): # ③ 近期不再用就切回分片
partition_again(layer.params) # 显存立刻归还
return x
重点看:只有「正在算的层 + 预取中的几层」占整块显存,这 就是 max_live_parameters 的存在理由(§5.3)。
3. 真实实现 I:参数怎么「一出生即分片」
3.1 zero.Init:猴子补丁式的构造期分片
deepspeed.zero.Init(partition_parameters.py:940)是个上下文管理器,继承自 InsertPostInitMethodToModuleSubClasses(:334)。进上下文时做两件事:
- 挂钩所有
nn.Module子类的构造:模块__init__一跑完,就对这个模块的每个参数执行分片——_partition_param(:1750)算出本 rank 的元素区间[rank*partition_size, (rank+1)*partition_size),拷进param.ds_tensor,然后free_param(:308)把param.data缩成 1 元素的占位。 - (可选)把
F.linear换成可处理分片权重的版本(zero3_linear_wrap)。
效果写在 Init 的 docstring 里:1T 参数模型,8 进程本来各要 4TB CPU 内存,分片后每个进程只承担 1/N,模型规模随聚合内存线性扩展。这也是「fp16 权重单卡放不下时必须用 zero.Init」的原 因(docstring 原话)。
3.2 不用 zero.Init 的事后补救
模型在上下文外构造好也没关系:DeepSpeedZeRoOffload.__init__(deepspeed/runtime/zero/parameter_offload.py:132)会调 _convert_to_zero_parameters(:252)——已有 zero 参数就用它把剩下的补转;一个都没有就现场开一个 Init(module=...) 把整模型转一遍。所以 ZeRO-3 对「先建模型后初始化」的用户代码同样成立,只是构造期内存峰值省不掉。
4. 真实实现 II:hook 机制——谁触发取/放
setup_zero_stage3_hooks(parameter_offload.py:293)→ _register_deepspeed_module(:331)递归给每个子模块挂上四个钩子:
| 钩子 | 触发时机 | 动作 |
|---|---|---|
| forward pre-hook | 模块前向前 | pre_sub_module_forward_function(:525):记 trace + fetch_sub_module(forward=True) |
| forward hook(后) | 模块前向后 | post_sub_module_forward_function(:549):release_sub_module(forward=True);顺带登记「悬空参数」(输出里出现的、属于外层模块的参数,:357-397) |
| forward hook(埋点) | 前向时插入 autograd 节点 | PreBackwardFunctionForModule(:460):backward 反传到该模块输出时,触发 pre-backward 取参 |
| forward pre-hook(埋点) | 同上 | PostBackwardFunctionModule(:491):模块所有输入的梯度都回流完时,触发 post-backward 释放 |
后两个是巧思所在:PyTorch 没有「模块级 backward 钩子」,于是 DeepSpeed 在前向图里插两个假的 autograd.Function(输入输出原样透传,只挂 backward),把「该模块的反向开始/结束」编码进 autograd 图本身。反向一到,钩子自然触发。ds_grads_remaining 引用计数处理「一个模块多个输入张量」的情况(最后一个梯度回来才算 backward 结束)。
四个函数都极其薄,全部动作转发给一个对象:PartitionedParameterCoordinator。
5. 真实实现 III:调度中枢 PartitionedParameterCoordinator
代码在 deepspeed/runtime/zero/partitioned_param_coordinator.py:73。三块机制:trace、fetch、release。
5.1 trace:跑一圈,把剧本记下来
生命周期是状态机(ZeRoTraceMode,:53):
第 0 步: INVALID ──(reset_step, :251)──► RECORD
第 1 步: RECORD ──(forward/backward 中 record_module, :222 逐模块记录)
──(reset_step: 构造参数级剧本 + 跨 rank 一致性断言)──► COMPLETE
第 2 步起: COMPLETE,__param_queue 按剧本供预取消耗
(发现实际顺序与剧本不符 → trace_prologue (:202) 整体作废重录)
reset_step(:251)里有三个 assert_ints_same_as_other_ranks——各 rank 的模块顺序、参数顺序、取用步数必须逐 rank 一致,否则 all-gather 会因集合通信错位而挂死。这就是为什么 ZeRO-3 讨厌数据依赖的动态分支(见 §7 坑)。