跳到主要内容

数据截至 (上游 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 = 128megablocks/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 范式按峰值预留的浪费。


3. 图示:稀疏矩阵长什么样

怎么读这张图: 横轴是拼起来的所有专家的 FFN 维,纵轴是补零后的 token 块。X 是非零块(要算),. 是零块(跳过)。每个行块恰好在一个专家的列范围内有一串非零块。

E0 的 ffn E1 的 ffn E2 的 ffn
┌───────────┐ ┌───────────┐ ┌───────────┐
行块0(E0) │ X X X X│ │ . . . .│ │ . . . .│
行块1(E0) │ X X X X│ │ . . . .│ │ . . . .│
行块2(E1) │ . . . .│ │ X X X X│ │ . . . .│
行块3(E2) │ . . . .│ │ . . . .│ │ X X X X│
行块4(E2) │ . . . .│ │ . . . .│ │ X X X X│
└───────────┘ └───────────┘ └───────────┘
每个行块的非零块数 = ffn/128(常数);哪个专家占哪几行,每步由路由决定

专家计算就拆成两个块稀疏 GEMM(SparseMLP.forwardmegablocks/layers/mlp.py:392-394):

  • stk.ops.sdd(x, w1.t(), topo):dense × dense,但只在稀疏块位置采样输出 → 稀疏的中间激活。
  • stk.ops.dsd(activation, w2):稀疏 × dense → 回到 dense 的 [padded_tokens, hs]

4. 真实实现:从路由输出到稀疏拓扑

dMoE 的前向在 sparse_forward_oncemegablocks/layers/dmoe.py:156-193)。拆成四步看。

4.1 第一步:元数据——比标准 MoE 多算一个「补零边界」

indices_and_padded_binsmegablocks/layers/dmoe.py:131-154)在第 1 章那套 sort / histogram / cumsum 之上,多算一组东西:

  • padded_tokens_per_expert = ops.round_up(tokens_per_expert, 128)——每个专家的 token 数向上取整到 128(megablocks/layers/dmoe.py:143-146round_up 本身是整除 trick,megablocks/ops/round_up.py:7-14)。
  • padded_bins = inclusive_cumsum(padded_tokens_per_expert)——补零后每个专家区间的结束位置。

于是同时握着两套边界:bins(真实 token 数,散射/加权用)和 padded_bins(补零后,矩阵形状用)。

4.2 第二步:padded_gather——把 token 搬进补零后的布局

ops.padded_gathermegablocks/backend/kernels.py:107-139)的输出有 padded_bins[-1] 行——形状是数据决定的,所以 wrapper 里有一次 padded_bins[-1].cpu().item() 的 host 同步(megablocks/backend/kernels.py:122-124)。这是整条管线里为数不多的同步点之一。

重排的核心是 Triton kernel _padded_copymegablocks/backend/kernels.py:45-104)。它的 grid 是「每个排序后的 (token, expert) 条目一个线程块」,索引数学就三行(megablocks/backend/kernels.py:60-76):

index_a = tl.load(indices + tl.program_id(0)) # 我在原始 (token×k) 展平序列里的位置
bin_idx = tl.load(bin_ids + tl.program_id(0)) # 我属于哪个专家
offset_in_bin = tl.program_id(0) - bins[bin_idx-1] # 我在本专家内部排第几
index_b = offset_in_bin + padded_bins[bin_idx-1] # 我在补零布局里的行号

方向标志 A_TO_B同一个 kernel 干两件事:gather 时 out[index_b] = x[index_a // TOP_K](注意除以 TOP_K——top-k 的多个条目映射回同一个源 token);scatter 时反向,并乘上路由权重。所以 padded_gather 的 backward 就是 padded_scattermegablocks/ops/padded_gather.py:29-42),padded_scatter 的 backward 就是 padded_gather,外加一个算路由权重梯度的 _padded_copy_wgradmegablocks/backend/kernels.py:228-276)。一对 kernel,四个用途。

4.3 第三步:topology——稀疏矩阵的「形状说明书」

topologymegablocks/layers/dmoe.py:68-129)生成 stk 需要的 CSR 风格元数据。妙处在于大部分元数据是免费的

  • offsets 不用算。 每个行块的非零块数恒为 ffn/128,所以行偏移就是等差数列,torch.arange 直接生成(megablocks/layers/dmoe.py:82-89)。
  • 只有 column_indices 要真算:行块 i 属于专家 j → 它的非零列块是 j × (ffn/128) + [0..ffn/128)。这由一个 CUDA 小 kernel 写出(ops.topologymegablocks/ops/topology.py:20-36ConstructIndicesKernelcsrc/indices.h:22-46)。注意输出 dtype 是 int16megablocks/ops/topology.py:28-33)——列块总数通常只有几千,16 位够用还省带宽。
  • data 是 meta 设备上的空壳。 拓扑只关心形状和索引,数值内存后面才分配,所以先用 device='meta' 占位省显存(megablocks/layers/dmoe.py:98-107,注释里作者自己说要清理这个设计)。

backward 还需要转置后的拓扑(算 dw1 时稀疏矩阵要转置)。sparse_transposemegablocks/layers/dmoe.py:35-66)的做法很巧:对 column_indices 再做一次 radix sort,用排序返回的 gather 顺序重排 row_indices,就得到了转置矩阵的列索引;转置的行偏移则是对列索引做 histogram + cumsum。排序位数同样按需截断(transpose_sort_end_bitmegablocks/layers/dmoe.py:28-33)。

4.4 第四步:专家 MLP

SparseMLP.forwardmegablocks/layers/mlp.py:379-394)最短形态就三行:sdd → act_fn → dsd。开 memory_optimized_mlp 则走手写 autograd 的 MemoryOptimizedMLPmegablocks/layers/mlp.py:187-306):backward 里重算激活函数而不是存它的输出,并且显式复用同一块显存放中间梯度(dactivation_fn_out = activation_fn_out 这类注释标明的复用,megablocks/layers/mlp.py:262-306)——省显存的思路与第 3 章 grouped 版本一脉相承。

4.5 原理演示

把整条链压成一段示意(示意,非源码):

# 示意,非源码
def dmoe_forward(x, router, w1, w2, block=128):
scores, weights, top_expert = router(x)
order = argsort(top_expert)
counts = bincount(top_expert, num_experts)
bins = cumsum(counts) # 真实边界
padded = round_up(counts, block) # 每专家补零到 128 的倍数
padded_bins = cumsum(padded) # 补零边界

Xp = zeros(padded_bins[-1], hidden) # 补零后的规整布局
for e in range(num_experts):
rows = order[bins[e-1]:bins[e]] # 专家 e 的真实 token
Xp[padded_bins[e-1] : padded_bins[e-1]+len(rows)] = x[rows]

topo = build_block_sparsity(padded_bins, block) # 每个 128 行块 → 一个专家的列块
H = sparse_sdd(Xp, w1.T, topo) # 只算非零块: [padded, E*ffn]
Y = sparse_dsd(gelu(H), w2) # [padded, hidden]

y = zeros_like(x)
for e in range(num_experts):
rows = order[bins[e-1]:bins[e]]
y[rows] += Y[padded_bins[e-1] : padded_bins[e-1]+len(rows)] * weights[rows]
return y # 每个 token 都算到了,零丢弃

5. 关键细节与坑

  • 128 是硬编码。 ffn_hidden_size 必须被 128 整除,否则 topology 直接 raise ValueError 并提示改配置(megablocks/layers/dmoe.py:71-75)。GLU 变体(SparseGLUmegablocks/layers/glu.py:18-61)共用同一拓扑。
  • 新 Triton 不可用。 mlp_impl='sparse' 在 triton ≥ 3.2.0 被 Arguments.__post_init__ 拒绝(megablocks/layers/arguments.py:75-85)——这也是 grouped 成为默认的现实原因,见第 3 章。
  • 两次 host 同步。 padded_bins[-1].cpu().item()megablocks/backend/kernels.py:122-124)和 GroupedMLPbatch_sizes.cpu()(第 3 章)都是变长输出的代价;作者能省则省,但没完全消掉。
  • 补零行的数值是零,散射时被天然丢弃。 散射只按 bins(真实边界)回写,padded_bins 只决定从哪读——所以补零行算了也白算,不会污染输出(kernel 索引数学保证,见 §4.2)。
  • 稀疏拓扑每个前向重建一次。 topologytorch.no_grad() 里、每个 forward 调一次(megablocks/layers/dmoe.py:177-178);几个小 kernel(sort/histogram/cumsum/indices)的开销就是「dropless」的固定税。

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

主题文件路径符号名
dMoE 前向megablocks/layers/dmoe.pyParallelDroplessMLP.sparse_forward_once
补零元数据megablocks/layers/dmoe.pyindices_and_padded_bins
稀疏拓扑megablocks/layers/dmoe.pytopologysparse_transpose
拓扑 CUDA kernelcsrc/indices.hConstructIndicesKernel
gather/scatter kernelmegablocks/backend/kernels.py_padded_copypadded_gatherpadded_scatter_padded_copy_wgrad
autograd 封装megablocks/ops/padded_gather.pymegablocks/ops/padded_scatter.pyPaddedGatherOpPaddedScatterOp
专家 MLP(sparse)megablocks/layers/mlp.pySparseMLP.forwardMemoryOptimizedMLP
向上取整megablocks/ops/round_up.pyround_up
GLU 变体megablocks/layers/glu.pySparseGLU