数据截至 (上游 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(_ParamAndGradBuffer,megatron/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 与四重区间映射
Range(megatron/core/optimizer/distrib_optimizer.py:78)就是个 (start, end),加一个 normalize 换基。真正的功夫在 _build_model_gbuf_param_range_map(distrib_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_ready(param_and_grad_buffer.py:863):只在最后一个 microbatch 且 overlap_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_grads(distrib_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-2896的copy_group_params)。
4.4 step 后:all-gather 补齐参数
参数方向的通信在 Bucket.start_param_sync(param_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_sync的force_all_reduce参数,:600;train_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.py | DistributedOptimizer |
| 区间类型 | megatron/core/optimizer/distrib_optimizer.py | Range |
| 四重区间映射 | 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 |
| 主分片写回 buffer | megatron/core/optimizer/distrib_optimizer.py | _copy_main_params_to_model_params |
| 连续梯度/参数 buffer | megatron/core/distributed/param_and_grad_buffer.py | _ParamAndGradBuffer |
| 梯度就绪 hook | megatron/core/distributed/param_and_grad_buffer.py | BucketGroup.register_grad_ready |
| reduce-scatter 发起 | megatron/core/distributed/param_and_grad_buffer.py | BucketGroup.start_grad_sync、finish_grad_sync |
| 参数 all-gather | megatron/core/distributed/param_and_grad_buffer.py | Bucket.start_param_sync、finish_param_sync |
| 优化器基类( 混合精度 step) | megatron/core/optimizer/optimizer.py | MixedPrecisionOptimizer |
下一章:05 · MoE 支持——专家并行怎么把「哪些 token 去找哪个专家」落成 All-to-All。