跳到主要内容

数据截至 (上游 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(LearnedRoutermegablocks/layers/router.py:61-113)。

难在哪? 难在「每个专家分到多少 token」是数据决定的,每一步都变。这一节展开讲。


2. 直觉:GPU 为什么讨厌变长

先看一个具象的例子。假设 batch 有 8 个 token、4 个专家、top-1 路由,某一步的分配结果是:

专家分到的 token 数
E03
E10
E24
E31

稠密 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 工程化的做法是:

  1. 给每个专家定一个容量 C = capacity_factor × (tokens × top_k / num_experts)——平均负载乘以一个富余系数。
  2. 每个专家分到容量就拒收:超出的 token 不过这个专家,输出里对应位置是零(靠残差连接兜底)——这就是 token dropping
  3. 不足容量的位置补零:反正形状固定了,多算一点无害的零。

这样每个专家的输入形状恒为 [C, hs],E 个专家拼成 [E, C, hs],可以用一次 torch.bmm 算完。问题解决了,代价是两个:

  • 丢 token 伤效果,且丢不丢、丢多少随每步的路由分布浮动。
  • 多出一个难调的超参 capacity_factor:小了丢得多,大了算得慢。

MegaBlocks 的标准 MoE 层就是这套范式的实现,下一节看它怎么落地。理解它,才能理解 dMoE 到底改掉了什么。


4. 真实实现:标准 MoE 的 binned 管线

标准 MoE 的专家计算在 ParallelMLPmegablocks/layers/moe.py:96)。它的前向(不开专家并行时)是 forward_oncemegablocks/layers/moe.py:209-235),整个管线分四步。

4.1 第一步:分拣元数据

indices_and_binsmegablocks/layers/moe.py:152-183)对拍平的 top_experts(长度 = tokens × top_k)做三个操作:

输出算什么用的算子
bin_ids, indices按专家 id 排序,得到每个位置的「原序号」ops.sort(cub radix sort)
tokens_per_expert每个专家分到几个 tokenops.histogram(cub histogram)
bins排序后每个专家的结束位置(前缀和)ops.inclusive_cumsum

两个值得记住的细节:

  • 排序只排低位。 key 是专家 id,64 个专家只要 6 位,于是 sort_end_bit = max(int(np.ceil(np.log2(num_experts))), 1)megablocks/layers/moe.py:110)传给 cub,不排整个 32 位(csrc/sort.h:23-60cub_radix_sortend_bit 参数)。
  • histogram 和 sort 是独立的两个 kernel。 代码里 TODO 注释在纠结「排好序的数据做 histogram 是不是更快,还是两个 kernel 并行更值」(megablocks/layers/moe.py:167-170)——看得出作者对每个小 kernel 的耗时都很在意。

4.2 第二步:算容量

expert_capacitymegablocks/layers/moe.py:133-137)就是 §3 的公式:

def expert_capacity(self, tokens: int) -> int:
world_size = mpu.get_expert_parallel_world_size(self.args)
tokens_per_expert = (self.top_k * tokens * world_size / self.num_experts)
return int(self.args.moe_capacity_factor * tokens_per_expert)

4.3 第三步:重排 → bmm → 还原

permute_and_computemegablocks/layers/moe.py:185-207)三行说完全部:

  • ops.binned_gather(x, indices, bins, expert_capacity, top_k):把 [tokens, hs] 按排序结果重排成 [E, C, hs] 的定容量桶。
  • self.mlp(x)MLP.forwardmegablocks/layers/mlp.py:162-167)就是 bmm → activation → bmm,权重是 [E, hs, ffn] 的三维张量,一次 bmm 算完所有专家。
  • ops.binned_scatter(...):乘上路由权重、散射回 [tokens, hs],top-k > 1 时沿 top-k 维求和(megablocks/backend/kernels.py:421-445)。

4.4 原理演示

把 §4.1~4.3 压缩成一段示意代码(示意,非源码):

# 示意,非源码
def moe_forward(x, router, expert_w1, expert_w2, capacity_factor):
scores, weights, top_expert = router(x) # 每个 token 挑 top-k 专家
order = argsort(top_expert) # 按专家 id 排序 → 重排索引
counts = bincount(top_expert, num_experts) # 每个专家几个 token
bins = cumsum(counts) # 排序后各专家的结束位置

C = int(capacity_factor * x.shape[0] / num_experts)
bucket = zeros(num_experts, C, hidden)
for e in range(num_experts):
rows = order[bins[e-1]:bins[e]] # 分给专家 e 的 token(排序后)
bucket[e, :len(rows)] = x[rows[:C]] # 超出 C 的部分被截断 = 丢 token

h = gelu(bmm(bucket, expert_w1)) # [E, C, ffn],一次 bmm
out = bmm(h, expert_w2) # [E, C, hs]

y = zeros_like(x)
for e in range(num_experts):
rows = order[bins[e-1]:bins[e]][:C]
y[rows] += out[e, :len(rows)] * weights[rows] # 加权散射回原位置
return y

重点看截断那一行:真实实现里「丢 token」不是一个显式的 if,而是固定形状的桶天然装不下

4.5 token 到底在哪一行被丢掉

binned 桶的填充在 Triton kernel _binned_copymegablocks/backend/kernels.py:326-389)里。它的 grid 是 (num_experts, expert_capacity)——每个线程块负责「专家 e 的第 i 个槽位」,关键三行(megablocks/backend/kernels.py:356-359):

# Calculate our offset into the input. If we don't
# have an input exit early.
if entry_idx >= num_tokens:
return
  • 当专家实际 token 数 少于容量:entry_idx >= num_tokens 的槽位直接早退,输出保持初始化的零——这是补零
  • 当实际 token 数 多于容量:grid 只开到 expert_capacity,序号 ≥ C 的 token 根本没有线程块来读它——这是丢弃。散射回去时它们的位置同样是零。

没有报错、没有标记,丢弃是完全静默的。这就是为什么论文和 README 都把「去掉 capacity_factor」当作卖点。

4.6 一个隐藏的逃生门

forward_once 里有这么一段(megablocks/layers/moe.py:218-223):

# If expert_capacity is set to zero, set the number of tokens
# per expert to the maximum we need to avoid dropping tokens.
expert_capacity = self.expert_capacity(sl * bs)
if expert_capacity == 0:
expert_capacity = torch.max(tokens_per_expert).item()

moe_capacity_factor 设为 0,容量就动态取「本步最热门专家的 token 数」——不丢 token 了,但容量变成数据依赖、且按最坏情况补齐。这说明作者很清楚两种范式的边界在哪:标准 MoE 的「定容量」和 dMoE 的「不丢」之间,其实只隔着这一行。


5. 关键细节与坑

  • 路由抖动(jitter)。 训练时可给路由输入乘 [1±eps] 的均匀噪声(LearnedRouter.jittermegablocks/layers/router.py:81-85),这是 Switch 留下的探索技巧,防止路由过早塌缩到少数专家。
  • 均匀分配开关。 uniform_expert_assignment=True 时用一个自定义 autograd Function 把专家 id 换成 arange % num_expertsmegablocks/layers/router.py:49-58)——纯为端到端 benchmark 用,注释明说「不收敛也能测性能」。
  • 桶是三维张量,显存按容量峰值占用。 [E, C, hs] 的大小由 capacity_factor 决定,与真实分布无关;分布越不均,浪费越多(inferred:这是定长方案的天然代价,代码里没有显式说明)。
  • top-k 求和的位置。 scatter 的输出先成 [tokens, top_k, hs]sum(dim=1)megablocks/backend/kernels.py:442-445),所以 top-k 越大,还原阶段的读写越重。

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

主题文件路径符号名
标准 MoE 专家计算megablocks/layers/moe.pyParallelMLPforward_oncepermute_and_compute
分拣元数据megablocks/layers/moe.pyindices_and_bins
容量公式megablocks/layers/moe.pyexpert_capacity
路由器megablocks/layers/router.pyLearnedRouter.forward_top_kjitter
定容量桶 kernelmegablocks/backend/kernels.py_binned_copybinned_gatherbinned_scatter
专家 MLP(bmm)megablocks/layers/mlp.pyMLP.forward
radix sort(cub)csrc/sort.hcub_radix_sortsort
histogram(cub)csrc/histogram.hcub_histogram
Python 侧 autograd 封装megablocks/ops/binned_gather.pyBinnedGatherOpBinnedScatterOp