跳到主要内容

数据截至 (上游 commit e79cb4c1bae1)

04 · 分布式优化器

这一章讲什么: Megatron 对「数据并行的显存与通信」的答案。DistributedOptimizer 把 Adam 状态按 DP rank 分片(ZeRO-1 形态),靠 _ParamAndGradBuffer 的「连续 buffer + 分桶 + hook」把梯度通信藏进反向过程。读完你会理解为什么 Megatron 的 DP 没有独立的梯度 all-reduce 步骤,以及「一个参数被两个 DP rank 各管一半」这种反直觉设计为什么是对的。


1. 它要解决的小问题

普通数据并行里,每张卡存全量的三样东西:

东西175B 模型的体量(量级)
参数(bf16)~350GB
梯度(bf16)~350GB
Adam 状态(fp32 主参数 + m + v)~1400GB

梯度通信本身(all-reduce)也是一步显眼的全局同步。ZeRO 的观察:DP 组里每张卡更新出的参数其实是一样的,那何必每卡都存全量优化器状态?切成 DP 份,每卡只更新 1/DP,更新完再互相补齐即可。

难点在工程:参数有成千上万、形状各异,「谁负责更新哪一段」怎么划分才不引入成堆的边界特判?通信又怎么塞回训练循环里不挡路?


2. 思路:切 buffer,不切参数

Megatron 的答案分两层,每层一个关键决策:

决策一:梯度先落进连续 buffer。 不按参数逐个同步,而是把同 dtype 的所有梯度拼成一根连续的一维 grad_data_ParamAndGradBuffermegatron/core/distributed/param_and_grad_buffer.py:1051),再按大小切成若干桶(bucket)

决策二:按区间归属,不按参数归属。 把每个 bucket 等分成 DP 份,DP rank r 拥有第 r 段连续区间——不管这段区间跨过几个参数、也不管一个参数是否被区间边界劈成两半。所有权只意味着「这段梯度的 reduce 归我收尾、这段参数的更新归我算」。

两个决策合起来的效果:

  • 反向时桶一满就异步 reduce-scatter——每卡只留下自己那段区间的梯度(通信和后续计算重叠);
  • step 时优化器只在自己的 fp32 主参数分片上跑 Adam;
  • step 后把新参数写回 buffer 的对应区间,all-gather 让所有 DP rank 拿到全量新参数。

图示:一个 step 里 DP 通信的两步

反向过程中 (按桶流水进行):
bucket 满了 ──► reduce-scatter ──► 每卡留下 1/DP 的梯度
▲ │
│ 与后续层的反向计算重叠 ▼
│ DistributedOptimizer.step():
│ 只在自己那段 fp32 主参数上跑 Adam
│ │
│ ▼
下一步前向前: 新参数写回 buffer 区间
all-gather ◄──────────────────────────┘
所有 DP rank 补齐全量参数

怎么读这张图: 上半是梯度方向(反向时发生,reduce-scatter),下半是参数方向(step 后发生,all-gather)。传统 DP 的「一次大 all-reduce」被拆成这两半,各自都能和计算重叠。


3. 原理演示:区间归属

这段演示「按段归属」为什么允许参数被劈开。关键:所有操作都是对一维 buffer 视图的切片操作,参数边界只出现在「映射表」里,不出现在通信里:

# 示意,非源码
buf = grad_buffer # 一根连续一维梯度, 例如 1000 个元素
dp_size = 4
shard = 1000 // 4 # 每 rank 拥有 250 个元素

my_range = Range(r*shard, (r+1)*shard) # 我拥有的区间, 如 [250, 500)
for param in params: # 参数可能跨区间边界!
overlap = intersect(param.range_in_buf, my_range)
if overlap:
# 我只管 param 落在 [250,500) 里的那一段
main_param_shard = fp32_copy(param)[overlap - param.start]
main_param_shard.grad = buf[overlap] # reduce-scatter 留给我的梯度

重点看:一个参数可以被两个 rank 各管一半,映射表(Range)负责记住「参数的第几段落在哪个 rank 的区间里」,通信和优化器全程只面对整齐的 1/DP 切片。


4. 真实实现

4.1 Range 与四重区间映射

Rangemegatron/core/optimizer/distrib_optimizer.py:78)就是个 (start, end),加一个 normalize 换基。真正的功夫在 _build_model_gbuf_param_range_mapdistrib_optimizer.py:134):为每个参数算四重区间,docstring(:154-160)写得明明白白:

区间名含义
gbuf_world参数在整个 grad buffer 里的全局区间
gbuf_world_in_bucket参数在所属 bucket 里的区间
gbuf_local参数在本 DP rank 本地视图里的区间
param分片在参数自身里的区间(参数被劈开时非零起点)

四者互转全靠 Range.normalize 平移。有了这张表,后续所有「模型梯度 → 主梯度」「主参数 → 模型参数」的搬运都是查表 + 切片。

_build_model_gbuf_range:200)负责先切出「本 rank 的区间」:bucket 大小必须能被 DP size 整除(有断言,:223-228),然后 gbuf_world_start = r * max_gbuf_range_size 逐 rank 排下去(:231-241)。注意它还保存所有 rank 的区间——reduce-scatter/all-gather 的参数构造要用(:204-208 的注释)。

4.2 反向途中:桶满即 reduce-scatter

梯度怎么进 buffer?DDP 包装给每个参数注册了梯度 hook,桶内所有参数的梯度就绪后触发同步。触发点是 BucketGroup.register_grad_readyparam_and_grad_buffer.py:863):只在最后一个 microbatchoverlap_grad_reduce=True 时计数(:874-876 的断言与注释),计数器与「黄金计数」相等就调 start_grad_sync:880-884)。

start_grad_sync:600)的关键分支在 :710-722:开了分布式优化器就不做 all-reduce,而是把 bucket.grad_data 按 DP 切出本地视图 local_data_view,调 dist_reduce_scatter_func(local_data_view, bucket.grad_data, ...)——每卡收走自己那 1/DP。多个桶还被 _coalescing_manager 合并成一次通信kernel(:706 的注释「Coalesce communication kernels across buckets」)。

4.3 step:只更新自己的分片

优化器本体是混合精度路线:模型参数 bf16,主参数 fp32 且只存分片。step 前后的两次搬运把 §2 的图落成了代码:

  • _copy_model_grads_to_main_gradsdistrib_optimizer.py:2765):docstring 第一句「Since this step follows a reduce-scatter through the DDP's grad buffer」——把 buffer 里自己那段梯度搬进主分片的 .grad。注意它用 param_range_map["param"] 切片(:2790-2794),正好处理「参数被劈开」的情形。
  • 之后基类跑 Adam(只在分片上)。
  • _copy_main_params_to_model_params:2811):docstring「Since this step is followed by an all-gather」——把更新后的主分片写回 buffer 的 gbuf_world_in_bucket 区间(:2874-2896copy_group_params)。

4.4 step 后:all-gather 补齐参数

参数方向的通信在 Bucket.start_param_syncparam_and_grad_buffer.py:383):对每桶 param_data 做 all-gather,overlap_param_gather=True 时同样异步发起(:407),在下一步前向真正用到该桶参数前 finish_param_sync:538)等待完成。至此一个 ZeRO-1 循环闭合:reduce-scatter(梯度) → 分片 step → all-gather(参数)。


5. 关键细节与坑

  • 这不是 ZeRO-3。 Megatron 这里分片的是优化器状态 + 梯度归宿,参数本体在 buffer 里仍是全量(all-gather 后每卡都有完整参数参与计算)。所以它对应 ZeRO-1(+梯度分片的 ZeRO-2 成分),显存大头省在 Adam 状态上。要更彻底的分片得走 Megatron-FSDP 路径(use_megatron_fsdp,在 _copy_model_grads_to_main_grads 里直接早退,distrib_optimizer.py:2776-2780)。
  • bucket 大小必须整除 DP size。 _build_model_gbuf_range 里有硬断言(:226-229);padding 由 buffer 布局层保证,用户一般不感知,但自定义参数分组时会踩到。
  • force_all_reduce 是逃生门。 要存全量 wgrad(如调试/检查点)时,训练循环会改走 all-reduce 而非 reduce-scatter(start_grad_syncforce_all_reduce 参数,:600train_step 里按 save_wgrads_interval 设置,megatron/training/training.py:3048)。
  • 梯度缩放因子在通信前施加。 gradient_scaling_factor != 1.0 时先原地乘(param_and_grad_buffer.py:667-672)——它同时承担「平均 vs 求和」和 MoE 模型的梯度修正(param_and_grad_buffer.py:1018-1020 的 docstring 说明)。顺序错了数值就错。
  • 多 DistOpt 实例要二次 all-reduce。 num_distributed_optimizer_instances > 1 时,reduce-scatter 只在实例内做,实例间还要补一次 all-reduce(param_and_grad_buffer.py:734-760)——大集群分层收敛的常见手法,但默认不开。
  • 异步通信的 handle 全局唯一。 grad_reduce_handle 断言「不应同时有多个未完成通信」(:632-634);这是靠「下一个桶组发起前先 drain 前驱桶组」(:621-631)维持的不变量,改这里极易引入竞态。

6. 代码地图

主题文件路径符号名
分布式优化器本体megatron/core/optimizer/distrib_optimizer.pyDistributedOptimizer
区间类型megatron/core/optimizer/distrib_optimizer.pyRange
四重区间映射megatron/core/optimizer/distrib_optimizer.py_build_model_gbuf_param_range_map_build_model_gbuf_range
梯度搬入主分片megatron/core/optimizer/distrib_optimizer.py_copy_model_grads_to_main_grads
主分片写回 buffermegatron/core/optimizer/distrib_optimizer.py_copy_main_params_to_model_params
连续梯度/参数 buffermegatron/core/distributed/param_and_grad_buffer.py_ParamAndGradBuffer
梯度就绪 hookmegatron/core/distributed/param_and_grad_buffer.pyBucketGroup.register_grad_ready
reduce-scatter 发起megatron/core/distributed/param_and_grad_buffer.pyBucketGroup.start_grad_syncfinish_grad_sync
参数 all-gathermegatron/core/distributed/param_and_grad_buffer.pyBucket.start_param_syncfinish_param_sync
优化器基类(混合精度 step)megatron/core/optimizer/optimizer.pyMixedPrecisionOptimizer

下一章:05 · MoE 支持——专家并行怎么把「哪些 token 去找哪个专家」落成 All-to-All。