数据截至 (上游 commit 952db33d6eac)
01 · MoE 为什么难算:变长专家输入
这一章讲什么: 先搞懂 MoE 层在算什么、为什么它对 GPU 这么不友好,再看业界(以及 MegaBlocks 的标准
MoE层)怎么用 capacity factor 和 token dropping 把问题摁成定长。这是理解第 2 章 dMoE 重构的绝对前提。
1. 它要解决的小问题
MoE 层在算什么? 一句话:把 Transformer 的 FFN 层复制 E 份(每份叫一个「专家」),再加一个「路由器」给每个 token 挑 top-k 个专家,只让这几个专家算这个 token,输出按路由分数加权求和。
为什么要这么做? 因为它把「参数量」和「计算量」解耦了:
- 参数量 ≈ 稠密 FFN 的 E 倍(E 个专家各有一份权重)。
- 计算量 ≈ 稠密 FFN 的 k 倍(每个 token 只过 k 个专家,k 通常是 1 或 2)。
这就是 MoE 吸引人的成本模型:花 1 倍的 FLOPs,买到 E 倍的参数。路由器的全部实现就是一个无 bias 的线性层加 softmax 加 top-k(LearnedRouter,megablocks/layers/router.py:61-113)。
难在哪? 难在「每个专家分到多少 token」是数据决定的,每一步都变。这一节展开讲。
2. 直觉:GPU 为什么讨厌变长
先看一个具象的例子。假设 batch 有 8 个 token、4 个专家、top-1 路由,某一步的分配结果是:
| 专家 | 分到的 token 数 |
|---|---|
| E0 | 3 |
| E1 | 0 |
| E2 | 4 |
| E3 | 1 |
稠密 FFN 是一次 [8, hs] @ [hs, ffn] 的大 GEMM。MoE 理想化地算,是 4 次小 GEMM:[3, hs] @ w1_0、[0, hs] @ w1_1、[4, hs] @ w1_2、[1, hs] @ w1_3。
问题出在这三个地方:
- 小。 单个专家的 token 数远小于整个 batch,GEMM 吃不满 GPU。
- 变。 m 维(token 数)每步都变,没法提前分配固定形状的 buffer,也没法用静态计算图优化。
- 偏。 有的专家可能是 0,有的可能是全部——负载完全不可控。
一句话直觉:GEMM 硬件喜欢「大而方」,路由给它的是「碎而变」。 后面所有的工程设计,都是在两者之间架桥。
3. 行业解法:capacity factor 与 token dropping
既然硬件要定长,那就把变长摁成定长。Switch Transformer / GShard 开创、Tutel 工程化的做法是:
- 给每个专家定一个容量 C =
capacity_factor × (tokens × top_k / num_experts)——平均负载乘以一个富余系数。 - 每个专家分到容量就拒收:超出的 token 不过这个专家,输出里对应位置是零(靠残差连接兜底)——这就是 token dropping。
- 不足容量的位置补零:反正形状固定了,多算一点无害的零。
这样每个专家的输入形状恒为 [C, hs],E 个专家拼成 [E, C, hs],可以用一次 torch.bmm 算完。问题解决了,代价是两个:
- 丢 token 伤效果,且丢不丢、丢多少随每步的路由分布浮动。
- 多出一个难调的超参
capacity_factor:小了丢得多,大了算得慢。
MegaBlocks 的标准 MoE 层就是这套范式的实现,下一节看它怎么落地。理解它,才能理解 dMoE 到底改掉了什么。