跳到主要内容

数据截至 (上游 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 里有两个形态,由 ZeroParamStatusdeepspeed/runtime/zero/partition_parameters.py:236)标记:

状态含义
AVAILABLE完整参数在卡上,可以参与计算
NOT_AVAILABLE只有 1/N 分片(在 param.ds_tensor 里),param.data 无意义
INFLIGHTall-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.Initpartition_parameters.py:940)是个上下文管理器,继承自 InsertPostInitMethodToModuleSubClasses:334)。进上下文时做两件事:

  1. 挂钩所有 nn.Module 子类的构造:模块 __init__ 一跑完,就对这个模块的每个参数执行分片——_partition_param:1750)算出本 rank 的元素区间 [rank*partition_size, (rank+1)*partition_size),拷进 param.ds_tensor,然后 free_param:308)把 param.data 缩成 1 元素的占位。
  2. (可选)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_hooksparameter_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 坑)。

5.2 fetch:取当前、等当前、顺手预取未来

fetch_sub_module:310)是核心,三步走(注释原文就是这三步):

  1. 发出当前模块参数的 all-gather__all_gather_params:569)。
  2. 等当前模块参数到位:每个参数从 InflightParamRegistry 取出 handle 等待;等待在独立的 __allgather_stream 上做,不阻塞主计算流的后续发射。
  3. 预取:trace 完整时,从 __param_queue 依序弹出后续参数,凑满 prefetch_bucket_size(默认 5e7 元素,zero/config.py:215)或顶到 max_live_parameters 上限为止,再发一批异步 all-gather。

all-gather 的本体在 all_gather_coalescedpartition_parameters.py:1512):把一组参数的分片按 ds_id 排序后拼成一条 flat buffer,一次集合通信取回整组_all_gather_dtype:1301),完成后各参数的 .data 指向 flat buffer 的对应视图。排序不是洁癖——集合通信要求各 rank 按同一顺序拼,注释里明说「顺序错了会静默拿到错误参数,极难调试」(:1550-1554)。

5.3 release:释放也要看复用距离

release_sub_module:494)把模块参数切回分片,但释放名单要过滤两道__params_to_release:634):

  1. 被预取标记过「马上还要用」的参数不放(避免预取白做);
  2. 从当前步向后扫 max_reuse_distance(默认 1e9 元素,zero/config.py:244)内的模块,扫到还会用的参数不放。

再加上两道总闸门:max_live_parameters(默认 1e9,zero/config.py:238)限制「整块形态」参数总量;param_persistence_threshold(默认 1e5,zero/config.py:221)让小参数直接常驻——太小的小参数参与 all-gather 得不偿失,mark_persistent_parametersparameter_offload.py:311)在建优化器时就把它们挑出来永不切分。

5.4 NVMe 参数的特殊预取

参数若 offload 在 NVMe 上(第 4 章),还要早一拍:__prefetch_nvme_param_partitions:667)扫剧本队列,把「排在在途参数之后」的 NVMe 分片先 swap_in 进 pinned 内存,等 all-gather 时数据已在 CPU,省掉磁盘等待。


6. 真实实现 IV:stage3 优化器——step 怎么只更新 1/N

DeepSpeedZeroOptimizer_Stage3deepspeed/runtime/zero/stage3.py:150)与协调器的分工:协调器管「参数整块/分片的形态」,优化器管「梯度与更新」

6.1 初始化:分片 + 分子组

  • 每个参数已有 ds_tensor 分片;优化器按 sub_group_size(默认 1e9 元素)把参数组再切成若干 sub-group_create_fp16_sub_groups:1160)——sub-group 是 step 和 offload 的调度单位。
  • _create_fp32_partitions:1031)为每个 sub-group 建本 rank 的 fp32 主权重分片;开 offload_optimizer 时这些分片直接放 CPU/NVMe(第 4 章)。
  • _setup_for_real_optimizer:712)建一条连续的 grad_partitions_flat_buffer:742):所有参数的梯度分片共享一块大内存,避免碎片化。

6.2 backward:梯度照样边算边分片

create_reduce_and_remove_grad_hooks:1437)给每个可训练参数注册 hook(z3 叶子模块的参数挂模块级 full backward hook,延迟到整模块反传完)。梯度就绪后进 IPG 桶(__add_grad_to_ipg_bucket:1543),桶满调 __reduce_and_partition_ipg_grads:1570),两条通信路径:

路径条件做法
__avg_scatter_contiguous_grads:1741连续梯度桶(默认)整桶 all_reduce 后,各 rank 只切走自己那段区间
__avg_scatter_grads:1791非连续reduce_scatter_coalesced 直接 scatter 到属主

拿到的梯度分片落进 grad_partitions_flat_buffer 的对应区间。全程没有任何时刻存在全量梯度。

6.3 step:逐 sub-group 轮转

step:2605)的主循环一行一行对应五个动作:

for sub_group_id, group in enumerate(self.fp16_groups):
_prepare_sub_group(sub_group_id) # 换入优化器状态/梯度(offload 时)
unscale_and_clip_grads(sub_group_id, global_norm) # 反缩放 + 裁剪
_optimizer_step(sub_group_id) # 真优化器只见本 rank 分片
_reassign_or_swap_out_partitioned_parameters(...) # fp32→fp16 写回分片(或换出 NVMe)
_release_sub_group(sub_group_id) # 释放/换出优化器状态

:2629-2641;各方法在 :2421:1185:2586:2449

_optimizer_step:1185)还是 stage 1/2 那一招:param_groups[group_id]['params'] = [fp32_param],真优化器只看见本 rank 的 fp32 分片。注意 step 完参数不会 all-gather 回整块——_pre_step:2365)先 _partition_all_parameters:3040)确保全是分片态,下一步 forward 再由协调器按需取回。这就是 stage 3 与 stage 1/2 最大的形态差异:参数永远不全,用时现取。


7. 关键细节与坑

  • 动态计算图是 ZeRO-3 的软肋。 trace_prologuepartitioned_param_coordinator.py:202)发现模块顺序与缓存不符就整本作废、退回无预取模式重录。每一两步作废一次 = 性能塌方。有数据依赖分支的模型要么绕开,要么接受无预取。
  • 取参通信藏不住时会直接暴露成 bubble。 预取深度由 prefetch_bucket_sizemax_live_parameters 共同决定;显存紧到 max_live_parameters 压得很低时,预取被节流,前向开始等通信。吞吐不对劲先查这两个值。
  • all-gather 顺序错了不会报错,只会算错。 all_gather_coalescedds_id 排序正是为此(partition_parameters.py:1512);safe_mode 下还有跨 rank 断言。自己往这套机制里塞自定义参数时,务必走同一套排序。
  • 叶子模块(z3 leaf)改变调度粒度。 z3_leaf_moduledeepspeed/utils/z3_leaf_module.py:15)标记的模块被当作整体取/放(不再递归子模块),配合 zero_module_granularity_threshold 可把巨型 kernel 融合模块当单原子调度,避免在模块内部碎取。
  • 悬空参数(dangling/external params)。 某模块输出了属于外层的参数张量时,forward hook 会把它登记到外层的 _external_paramsregister_external_parameterpartition_parameters.py:150),防止它被内层提前释放。HuggingFace 模型里 logits/past_key_values 一类输出常见这种情况。
  • backward hook 埋点对「输出为 None / dict / 自定义对象」敏感。 _post_forward_module_hook 里有专门的分支枚举输出里的张量(parameter_offload.py:357-386),识别不了的输出类型会让 post-backward 钩子不触发(代码里留了 TODO 注释,:497-503 附近)。表现是显存涨——参数没被释放,不是算错。
  • 跨 rank 一致性是硬要求。 reset_step 里三处 assert_ints_same_as_other_rankspartitioned_param_coordinator.py:260-263):任何「各 rank 模块/参数顺序不同」的模型结构在 ZeRO-3 下会直接挂。

8. 代码地图

主题文件路径符号名
参数状态机deepspeed/runtime/zero/partition_parameters.py:236ZeroParamStatus
构造期分片deepspeed/runtime/zero/partition_parameters.py:940:334InitInsertPostInitMethodToModuleSubClasses
单参数分片/释放deepspeed/runtime/zero/partition_parameters.py:1750:308_partition_paramfree_param
协聚 all-gatherdeepspeed/runtime/zero/partition_parameters.py:1512:1301all_gather_coalesced_all_gather_dtype
hook 装配deepspeed/runtime/zero/parameter_offload.py:293:331setup_zero_stage3_hooks_register_deepspeed_module
backward 埋点deepspeed/runtime/zero/parameter_offload.py:460:491PreBackwardFunctionForModulePostBackwardFunctionModule
前/后向四个薄函数deepspeed/runtime/zero/parameter_offload.py:525-603pre/post_sub_module_forward/backward_function
调度中枢deepspeed/runtime/zero/partitioned_param_coordinator.py:73PartitionedParameterCoordinator
取/预取deepspeed/runtime/zero/partitioned_param_coordinator.py:310fetch_sub_module__all_gather_params_
放/复用窗口deepspeed/runtime/zero/partitioned_param_coordinator.py:494:634release_sub_module__params_to_release
trace 生命周期deepspeed/runtime/zero/partitioned_param_coordinator.py:251reset_steprecord_moduletrace_prologue
stage3 优化器deepspeed/runtime/zero/stage3.py:150DeepSpeedZeroOptimizer_Stage3
step 轮转deepspeed/runtime/zero/stage3.py:2611step_prepare_sub_group_optimizer_step_reassign_or_swap_out_partitioned_parameters
梯度分片通信deepspeed/runtime/zero/stage3.py:1745:1791__avg_scatter_contiguous_grads__avg_scatter_grads