数据截至 (上游 commit 952db33d6eac)
02 · dropless MoE 与 block-sparse 重构
这一章讲什么: MegaBlocks 的立身之本——怎么做到「一个 token 都不丢」还不亏性能。读完你会理解:补零到 128 的倍数为什么这么重要、block-sparse 矩阵的拓扑怎么一步步算出来、以及
padded_gather/padded_scatter这对 Triton kernel 的索引数学。
1. 它要解决的小问题
第 1 章的结论是:标准 MoE 为了讨好 GEMM,给每个专家定死容量、丢多补少。dMoE 要回答的问题是:能不能一个 token 都不丢,还让 GEMM 高兴?
约束摆得很清楚:
- 不丢 → 每个专家实际分到多少 token 就得算多少,m 维仍然变。
- 要快 → 不能退化成「每专家一个小 GEMM」。
- 要少补 → 补零是白算的 FLOPs,补得越少越好。
2. 思路:把「变长」变成「块级规整」
dMoE 的让步很小:不追求定长,只追求「块级定长」——把每个专家的 token 数向上取整到 128 的倍数(self.blocking = 128,megablocks/layers/dmoe.py:24)。
这一步带来两个结构性后果:
- 每个 128-token 的行块恰好属于一个专家。 因为每个专家的区间都是 128 的整数倍,128 对齐切分后,不会有一个块跨两个专家。
- 于是整个计算可以看成一张块稀疏矩阵乘。 把中间激活想成
[padded_tokens, E × ffn]的大矩阵:token 块 i 只属于专家 j,所以它只有专家 j 的那ffn/128个列块非零,其余全是零——而这些零根本不用算。
稀疏 pattern 每步都变(取决于路由),但块的大小恒为 128×128——这正是块稀疏 GEMM 库(stanford-stk)能吃的东西。变长问题就此消失,只剩「稀疏 pattern 是什么」这个纯元数据问题。
补零的代价有多小?每个专家最多多算 127 个空 token;实验规模下通常远小于 capacity_factor 范式按峰值预留的浪费。