MoE 稀疏路由的工程真相:负载均衡、容量因子与通信瓶颈
混合专家的论文读起来非常清爽:一个路由器算 softmax,取 top-k,加一个负载均衡的辅助损失,结束。真正在 8×A100 上把它跑起来之后才发现,论文省略的那部分才是工程量所在。这篇记录一次从「能跑」到「跑得住」的完整调参过程。
一、稀疏性带来的到底是什么
先把账算清楚。假设一个 dense 模型每层 FFN 的参数量是 P,一次前向的 FLOPs 大致是 2P(乘加各一次)× token 数。换成 E 个专家、每个 token 激活 k 个:
- 总参数量变成 E·P,模型容量提升 E 倍
- 每个 token 的计算量是 k·2P,只增长 k 倍
- 取 E = 64、k = 2,就是参数量 ×64、计算量 ×2
看上去是白捡的便宜。代价藏在两个地方:显存装得下这些参数吗,以及把 token 送到对应专家所在的设备上要花多少通信。
二、路由器只是一个线性层
令人意外的是,路由器几乎总是一个不带偏置的单层线性映射,加不加激活、加不加归一化,效果差别都在噪声范围内。
class Router(nn.Module):
def __init__(self, d_model, n_experts):
super().__init__()
self.w = nn.Linear(d_model, n_experts, bias=False)
def forward(self, x): # x: [tokens, d_model]
logits = self.w(x.float()) # 关键:路由永远用 fp32
probs = logits.softmax(dim=-1)
w, idx = probs.topk(self.k, dim=-1)
w = w / w.sum(dim=-1, keepdim=True)
return w, idx, probs
注释里那行是花了两天才定位的问题。BF16 下路由 logits 的舍入误差足以让 top-k 的选择在前后两次相同输入上不一致,进而导致梯度累积阶段的专家分配漂移,训练损失出现周期性尖峰。把路由部分强制提到 fp32 之后曲线立刻平了。
三、负载均衡:辅助损失怎么写
如果不加约束,路由器会迅速塌缩到只用少数几个专家——这是一个正反馈:某个专家因为随机初始化稍微好一点,被选中更多,训练得更好,于是被选中更多。
标准做法是加一个辅助损失,惩罚"分配比例"与"路由概率"的相关性:
其中 fe 是实际被分到专家 e 的 token 比例,Pe 是该批次里路由器给 e 的平均概率。两者都均匀时,Laux = 1,取到最小值。
这个损失的权重 α 很敏感,我们扫过的结果:
| α | 专家利用率(变异系数) | 验证损失 | 备注 |
|---|---|---|---|
| 0 | 1.41 | 2.38 | 7/64 专家承担 60% 流量 |
| 0.001 | 0.62 | 2.21 | 仍有明显长尾 |
| 0.01 | 0.18 | 2.14 | 采用 |
| 0.1 | 0.05 | 2.29 | 均衡了,但路由退化成近似随机 |
α = 0.1 那一行值得注意:负载完全均衡反而变差,说明均衡本身不是目标。理想状态是路由器按内容做出有区分度的选择,同时统计上不至于饿死大部分专家。α 的作用是给这个平衡加一个软约束,不是硬性拉平。
四、容量因子与丢 token
专家是在固定形状的张量上批量计算的,所以每个专家有一个容量上限:
超出容量的 token 会被直接丢弃——它们的 FFN 输出为零,只保留残差连接。这在训练时是可接受的正则化,在推理时则是明确的质量损失。
我们在 4096 token/批、64 专家、top-2 的配置下测了不同容量因子:
| capacity_factor | 训练丢弃率 | 激活显存 | 单步耗时 |
|---|---|---|---|
| 1.0 | 11.3% | 基准 | 基准 |
| 1.25 | 3.8% | +18% | +9% |
| 1.5 | 1.1% | +37% | +21% |
| 2.0 | 0.2% | +82% | +44% |
最终选了训练 1.25、推理 2.0 的组合。推理阶段批次小、波动大,容量必须留足;训练阶段少量丢弃对最终指标影响不到 0.02,不值得为此付 40% 的时间。
五、真正的瓶颈是 all-to-all
专家并行的意思是把 E 个专家切分到 N 张卡上,每张卡持有 E/N 个。于是每一层都要做两次 all-to-all:一次把 token 按目标专家分发出去,一次把结果收回来。
在我们的 8 卡 NVLink 机器上,单层 all-to-all 的耗时约占整层的 28%;跨机(200 Gbps IB)时这个数字涨到 61%。几个有效的缓解手段:
- 先做本地 top-1 聚合:如果 top-2 的两个专家恰好在同一张卡上,合并成一次发送。实测能减少 12%–15% 的通信量。
- 通信与计算重叠:把注意力层的计算和上一层 MoE 的 all-to-all 放进不同 stream。这一项收益最大,端到端提速 17%。
- 专家分组放置:统计训练中期的共现矩阵,把经常被同一批 token 同时选中的专家放到同一节点。收益约 6%,实现复杂度不低,视情况取舍。
- 把 token 按目标专家排序后再发:避免分散的小消息,让 NCCL 走大块传输路径。
六、能不能不用 MoE
写完这些之后必须说句公道话:如果目标是固定显存预算下的最佳质量,dense 模型在很多规模上仍然更划算。MoE 的优势场景是显存充裕但算力受限——例如推理侧要服务大量并发、算力是瓶颈,或者训练侧希望在相同 FLOPs 预算下拿到更大的模型容量。
如果部署环境是单卡消费级显卡,一个 7B 的 dense 模型几乎一定优于总参数 40B、激活 7B 的 MoE,因为后者根本装不下。选型时先问显存还是算力更紧张,答案往往就出来了。