跳到主要内容

数据截至 (上游 commit 952db33d6eac)

04 · 路由、负载均衡与专家并行

这一章讲什么: 前三章解决了「怎么算得快」,这章解决「怎么分得均」和「怎么摆到多卡」。路由器和负载均衡决定 MoE 训得好不好;parallel_forward_once 的三次置换 + all_to_all 决定 MoE 能不能扩展到几十上百卡。最后看与 Megatron-LM 的接合面。


1. 它要解决的小问题

拆开是两个问题:

  • 算法面: 路由器是学出来的。如果它把 90% 的 token 都发给同一个专家,那个专家撑爆、其余专家荒废——需要负载均衡损失把分配「掰匀」。
  • 系统面: 64 个专家的权重可能一张卡放不下,要把专家切到多卡;可 token 在每张卡的本地 batch 里,得有一种高效方式把 token 送到「拥有它目标专家」的那张卡。

2. 路由器:一个线性层的全部戏份

LearnedRouter.forwardmegablocks/layers/router.py:92-113)短短二十行,但每一行都是一个经典 MoE 技巧:

步骤干什么细节
jitter训练时给输入乘 [1-eps, 1+eps] 均匀噪声megablocks/layers/router.py:93-94:81-85
logits无 bias Linear:[tokens, hs] → [tokens, E]权重每卡一份,不并行(:66-69
存 logits追加到全局 _ROUTER_LOGITS,供 z-loss_save_router_logits:13-17
softmax → top-k取分数与专家 id;k=1 走 max 快路径:98-99_top_k :87-90
权重归一化可选按 p 范数归一 top-k 权重moe_normalize_expert_weights:101-106
benchmark 开关可强制均匀分配(只测性能不收敛)_UniformExpertAssignment:49-58

返回三件套:scores(全 E 个专家的 softmax 分,负载均衡损失要用)、expert_weights(top-k 分)、top_experts(top-k 专家 id)。

z-lossbatched_router_zlossmegablocks/layers/router.py:25-46):对每个路由器的 logits 算 mean(logsumexp(logits)²),惩罚过大的 logits——这是 ST-MoE 提出的稳定化手段。同样走「各层存全局列表、宿主统一结算」的模式。


3. 负载均衡损失:全局列表 + 一次结算

每层该算的部分很小(ParallelMLP.load_balancing_lossmegablocks/layers/moe.py:138-150):

scale = self.num_experts / (tokens * self.top_k)
return scale * torch.dot(tokens_per_expert, expert_scores.mean(dim=0))

含义:实际 token 分布(tokens_per_expert)与平均路由概率(expert_scores.mean)的点积——分布越偏,loss 越大。这是个可微的代理:梯度只通过 expert_scores 流向路由器。

值得注意的是它不在这里 backwardParallelMLP.forward 只做一件事:训练时把 (tokens_per_expert, scores) append 到模块级全局列表(megablocks/layers/moe.py:430-431save_load_balancing_loss :17-19)。真正结算在 batched_load_balancing_lossmegablocks/layers/moe.py:32-89),由训练宿主在每个 step 调用:

  • num_layers / pipeline_model_parallel_size 校验本 stage 该有几条记录(megablocks/layers/moe.py:39-57)——这是给 Megatron 流水线并行设计的。
  • 把所有层的 tokens_per_expertexpert_scores 拼起来做一次点积,总缩放是 loss_weight × num_experts / (num_layers × tokens × top_k)megablocks/layers/moe.py:86-91)。
  • 用完要 clear_load_balancing_loss()megablocks/layers/moe.py:27-29)。

设计要点:库只管「记账」,不管「入账」。 这是「层实现而非框架」定位的典型体现——loss 怎么加进总目标、什么时候清零,全由 Megatron 侧的训练循环决定。


4. 专家并行:三次置换 + all_to_all

moe_expert_model_parallelism=True 后,forward_fn 切到 parallel_forward_oncemegablocks/layers/moe.py:129-130:237-423)。函数开头的注释把算法讲得很白(megablocks/layers/moe.py:238-256):

  1. 本地置换:把本地 token 按目标专家排序——这样发给同一张卡的 token 是连续的。
  2. 跨卡置换(all_to_all):每张卡拿到「属于自己专家」的全部 token。
  3. 再本地置换:all_to_all 后 token 按来源卡分组,不是按专家分组,要重排一次才能算。
  4. 算完 MLP 后原路倒放:再 all_to_all 回去、散射还原。

4.1 图示

怎么读这张图: 以 2 卡、4 专家为例,从左到右是 token 的三次重排。

rank0 本地 ①按专家排序 ②all_to_all ③按专家重排
[t0→E2 t1→E0] → [E0:t1 E2:t0] → rank0 收 E0,E1 → [E0:... E1:...]
[t2→E1 t3→E3] [E1:t2 E3:t3] rank1 收 E2,E3 [E2:... E3:...]
(按来源卡分段) (按专家分段)


④ 本地专家 MLP → 逆序送回

4.2 真实实现里的四个动作

动作一:先交换 token 计数。 在搬数据之前,先用异步 dist.all_to_all_singletokens_per_expert 本身交换一轮(megablocks/layers/moe.py:275-281)——每张卡得先知道会收到多少 token,才能分配 buffer。async_op=True 让它和下面的本地 gather 重叠。

动作二:本地 gather + 算 send/recv。 ops.gather 把 token 按专家聚拢(megablocks/layers/moe.py:286-290),然后把计数搬上 CPU 求和得到 send_counts / recv_countsmegablocks/layers/moe.py:304-313)——all_to_all_single 的 split sizes 必须是 host 侧的 list,这又是一次必要的 host 同步。

动作三:异步 all_to_all + 重建专家分组。 数据在路上(megablocks/layers/moe.py:323-329)的同时,CPU/GPU 并行准备第二次本地置换的元数据:

  • parallel_tokens_per_expert 重建每个专家在本卡的边界(replicate_binsmegablocks/layers/moe.py:339-346);
  • ops.replicate 按计数把专家 id「铺」成与收到 token 等长的数组(megablocks/layers/moe.py:348-358);
  • 再排一次 radix sort 得到按专家分组的顺序(megablocks/layers/moe.py:361-365;旁边 TODO 说这里的 sort_end_bit 还能再缩小)。

动作四:算完原路返回。 permute_and_compute 算完后,第二次 all_to_all 送回(megablocks/layers/moe.py:406-411),若开了隐藏维切分先 ops.sum 归并(megablocks/layers/moe.py:413-419),最后 ops.scatter 散射还原(megablocks/layers/moe.py:422)。

all_to_all 本身是可微的:AllToAllOpmegablocks/layers/all_to_all.py:8-44)的 backward 就是把 input/output split sizes 互换再来一次 all_to_all。

4.3 专家怎么切:专家维 × 隐藏维

mpu 定义了二维切分(megablocks/layers/mpu.py:65-93):

  • expert_sharding_degree = min(world_size, num_experts):先按专家个数切。
  • hidden_sharding_degree = world_size / esd:卡比专家多时,剩余维度切 FFN 的隐藏维——相当于专家内部的 tensor parallelism。

初始化保证两种切法结果一致:create_moe_expert_weightsmegablocks/layers/mlp.py:43-88每个 rank 都先建完整主权重、同一种子初始化、再切出自己那片——注释明说「同样的 seed 在数据并行和专家并行下采到同样的权重」(megablocks/layers/mlp.py:49-51)。

还有一个容易漏的细节:开专家并行时权重梯度要乘 1 / world_sizeScaleGradientmegablocks/layers/mlp.py:18-32:148-151)——因为同一专家的梯度会在专家并行组内被梯度同步再累加一次,提前缩放抵消重复计数(inferred:代码只给缩放值,未写原因)。


5. 与 Megatron-LM 的接合面

代码里能看到三条明确的接合线:

接合点形态位置
配置from_megatron 把 Megatron args 按同名字段拷成 Argumentsmegablocks/layers/arguments.py:95-99
参数标记给专家权重打 expert_model_parallel 属性,供优化器/权重衰减分组识别megablocks/layers/mpu.py:31-41
FSDP 提示注释要求 FSDP 包住 ParallelMLP,让权重 all-gather 排在专家 all2all 之前megablocks/layers/moe.py:92-95

训练入口不在本库:exp/dmoe/dmoe_125m_8gpu.sh 等脚本本质是 Megatron 的 pretrain_gpt.py--moe-num-experts 等参数。流水线并行的层数校验(§3)也是为 Megatron 的 stage 划分写的。


6. 关键细节与坑

  • 路由器是每卡一份的冗余。 注释明说路由权重不做专家并行,因为每卡都要路由自己的 batch(megablocks/layers/router.py:66-69)。专家数极大时,这份冗余和它的 logits 存储都要算进显存账。
  • host 同步点比单机路径多。 token 计数 all_to_all 后要上 CPU 算 split sizes(megablocks/layers/moe.py:304-313),代码里 TODO 说也许该在 GPU 上算完再传回 host。
  • 流水线下 loss 记录数必须对上。 batched_load_balancing_loss 校验记录数等于本 stage 层数,对不上直接 raise(megablocks/layers/moe.py:44-57)——忘了清零或层数配错会在这一步炸出来。
  • shared expert 是叠加不是替代。 开了 shared_expert,稠密 MLP 的输出与专家输出相加(可选按 top_k+1 加权,megablocks/layers/mlp.py:554-571megablocks/layers/moe.py:468-473)——这是 DeepSeek 系流行的共享专家设计,但它走稠密路径,与本文的稀疏主线正交。

7. 代码地图(本章涉及)

主题文件路径符号名
路由器megablocks/layers/router.pyLearnedRouter.forwardjitter_top_k
z-lossmegablocks/layers/router.pybatched_router_zloss_save_router_logits
负载均衡损失megablocks/layers/moe.pyload_balancing_lossbatched_load_balancing_losssave_load_balancing_loss
专家并行前向megablocks/layers/moe.pyParallelMLP.parallel_forward_once
all_to_all autogradmegablocks/layers/all_to_all.pyAllToAllOp
并行拓扑megablocks/layers/mpu.pyexpert_sharding_degreehidden_sharding_degreeexperts_per_rank
初始化一致性megablocks/layers/mlp.pycreate_moe_expert_weights
梯度缩放megablocks/layers/mlp.pyScaleGradientscale_gradient
Megatron 适配megablocks/layers/arguments.pyfrom_megatronArguments
共享专家megablocks/layers/mlp.pySharedMLPadd_experts_sharedexpert