数据截至 (上游 commit 952db33d6eac)
03 · grouped GEMM 路径
这一章讲什么: dMoE 的第二种专家计算后端。第 2 章用「补零 + 块稀疏矩阵」消灭变长;这一章用「分组 GEMM」直接接受变长。读完你会知道
gmm的调用约定、为什么要手写 autograd、以及 sparse / grouped 两条路径怎么选。
1. 它要解决的小问题
第 2 章的 block-sparse 方案已经很好,但仍有两根刺:
- 补零是白算的 FLOPs。 每个专家最多浪费 127 个 token 的计算——不多,但不是零。
- stk 这条链脆弱。 sparse 后端依赖 stanford-stk,且在 triton ≥ 3.2.0 直接不可用(
megablocks/layers/arguments.py:75-85)。
于是就有一个自然的问题:能不能不补零,让 GEMM 库直接接受「每个专家 m 不同」的输入? grouped GEMM 就是干这个的。
2. 思路:把变长原样交给分组 GEMM
grouped GEMM 的约定:输入是所有组按行连续拼起来的大矩阵 [sum(m_i), k],权重是 [num_groups, k, n],再给一个 batch_sizes = [m_0, m_1, ...] 说明每组的行数。一次 kernel 调用算完所有组,各组 m 可以不同。
对照 MegaBlocks 的问题,这简直是量身定做:
- 「组」= 专家,
batch_sizes=tokens_per_expert——第 1 章的 histogram 已经算好了。 - 「按行连续拼接」= 按专家排序后的 token 布局——第 1 章的 sort 已经给好了。
- 于是 gather/scatter 退化成无补零版本(
ops.gather/ops.scatter,内部就是padded_bins == bins的特例,megablocks/backend/kernels.py:141-170、:206-208)。
block-sparse 方案里的「补零 + 拓扑 + sdd/dsd」三件套,整体被一句 gmm 换掉。
3. 图示:两条路径的分叉
怎么读这张图: 上半是第 2 章的 sparse 路径,下半是本章的 grouped 路径。分叉点在「重排之后」;grouped 少了补零和拓扑两个环节。
sparse: gather(补零) → [padded_tokens] → 建拓扑 → sdd → act → dsd → scatter
↑ 依赖 stk / triton 版本
grouped: gather(不补零) → [tokens] ─────────────→ gmm → act → gmm → scatter
↑ batch_sizes = tokens_per_expert
代码上就是 ParallelDroplessMLP.forward_once 的一个 if(megablocks/layers/dmoe.py:282-286):mlp_impl == 'sparse' 走 sparse_forward_once,否则走 grouped_forward_once(megablocks/layers/dmoe.py:239-258)。
4. 真实实现
4.1 前向:元数据都回到「无补零」版本
grouped_forward_once(megablocks/layers/dmoe.py:239-258)调用的是第 1 章的 indices_and_bins(不算 padded_bins),然后 grouped_permute_and_compute(megablocks/layers/dmoe.py:260-280)三步:
ops.gather(x, indices, bin_ids, bins, top_k)——按排序结果聚拢,无补零(megablocks/layers/dmoe.py:273-274)。self.mlp(x, tokens_per_expert)——GroupedMLP,下一节看。ops.scatter(...)——加权散射回原位置,top-k 求和。
4.2 GroupedMLP:三行 gmm 与一次 host 同步
GroupedMLP.forward(megablocks/layers/mlp.py:501-523):
def forward(self, x, tokens_per_expert):
batch_sizes = tokens_per_expert.cpu().to(torch.long)
...
x = gg.ops.gmm(x, w1, batch_sizes, trans_b=True)
x = self.args.activation_fn(x)
return gg.ops.gmm(x, w2, batch_sizes)
两个要点:
batch_sizes必须在 CPU。 第一行的.cpu()是一次 host 同步——grouped_gemm 的 kernel 调度要在 host 侧知道每组大小。这是 grouped 路径最显眼的同步开销。- 权重 view 成三维。
SparseMLP存的w1是[E*ffn, hs]的二维张量,这里view(ne, -1, hidden_size)成[E, ffn, hs]喂给 gmm(megablocks/layers/mlp.py:506-509)。所以 GroupedMLP 直接继承 SparseMLP 的初始化和权重布局(megablocks/layers/mlp.py:499)。
gg 是 grouped_gemm_util(megablocks/grouped_gemm_util.py:5-26):import 不到 grouped_gemm 包就 warn,Arguments.__post_init__ 在 mlp_impl='grouped' 时硬性 assert 可用(megablocks/layers/arguments.py:87-88)。依赖钉在 grouped_gemm==0.3.0(setup.py:72)。
4.3 手写 autograd:激活重算 + 显存复用
朴素版 forward 要存两个中间量(gmm 输出、激活输出)。MemoryOptimizedGroupedMLP(megablocks/layers/mlp.py:397-496)用手写 torch.autograd.Function 把显存压下去,三招:
| 招 | 做法 | 位置 |
|---|---|---|
| 激活重算 | backward 里重放 activation_fn(sdd_out),用 PyTorch autograd 顺手拿激活的梯度函数,而不是存激活输出 | megablocks/layers/mlp.py:447-452 |
| 原位复用 | gmm(..., c=dactivation_fn_out)、gmm(..., c=ddsd_out)——把梯度直接写进不再需要的中间量显存 | megablocks/layers/mlp.py:465-473、:490-493 |
| 融合激活 backward | 默认 GELU 时用融合的 gelu_backward_ 一步到位 | megablocks/layers/mlp.py:476-480、megablocks/layers/gelu.py |
MemoryOptimizedMLP(sparse 版,megablocks/layers/mlp.py:187-306)是同一思路在 stk 算子上的版本;GLU 也有对应的 MemoryOptimizedGroupedGLU(megablocks/layers/glu.py:71-170),但 sparse GLU 的 memory_optimized 直接 NotImplementedError(megablocks/layers/glu.py:47-52)。
4.4 原理演示
示意,非源码:
# 示意,非源码
def grouped_mlp(x, w1, w2, tokens_per_expert):
# x: [sum(tokens_per_expert), hs],各专家段连续排列
h = grouped_gemm(x, w1, batch_sizes=tokens_per_expert) # 每组 m 不同
h = gelu(h)
return grouped_gemm(h, w2, batch_sizes=tokens_per_expert)
class RecomputeMLP(torch.autograd.Function):
@staticmethod
def forward(ctx, x, w1, w2, sizes):
h = grouped_gemm(x, w1, sizes)
ctx.save_for_backward(x, w1, w2, sizes, h) # 只存 gmm 输出,不存激活输出
return grouped_gemm(gelu(h), w2, sizes)
@staticmethod
def backward(ctx, dY):
x, w1, w2, sizes, h = ctx.saved_tensors
a = gelu(h) # 重算激活
da = dY_mm_w2 # 复用 dY 的显存放回传梯度
dh = da * gelu_grad(h)
return dh_mm_w1, dhT_mm_x, aT_mm_dY, None
4.5 两条路径怎么选
| 维度 | sparse(第 2 章) | grouped(本章) |
|---|---|---|
| 补零开销 | 每专家最多 127 token | 零 |
| 依赖 | stanford-stk,triton < 3.2.0 | grouped_gemm 0.3.0 |
| 拓扑计算 | 每前向建稀疏拓扑 + 转置 | 不需要 |
| host 同步 | padded_bins[-1].cpu() | batch_sizes.cpu() |
| 官方定位 | 老路径 | 默认 + Hopper 推荐(README) |
默认配置已经回答了这个问题:mlp_impl: str = 'grouped'(megablocks/layers/arguments.py:50)。注册表在 dmlp_registry(megablocks/layers/dmlp_registry.py:12-17),mlp_type(mlp/glu)× mlp_impl(sparse/grouped)四种组合都能选。
5. 关键细节与坑
- 专家并行时特意复用 CPU 张量。
parallel_forward_once里 grouped 分支直接把上一段通信时已搬到 CPU 的parallel_tokens_per_expert_cpu求和复用,注释明说「避免多一次设备 同步」(megablocks/layers/moe.py:383-390)。能看出作者对 host 同步点很敏感。 - gather/scatter 与补零版共享 kernel。 无补零就是把
padded_bins传成bins(megablocks/backend/kernels.py:206-208的scatter),所以第 2 章的索引数学这里完全适用。 - 显存优化版是训练默认假设。 测试里
memory_optimized_mlp=True(tests/layers/dmoe_test.py:73);关掉它就走朴素 gmm 三行。 - 性能数字看 benchmark。 仓库自带
megablocks/ops/matmul_benchmark.py和permute_benchmark.py等微基准,论文的 40% / 2.4x 数字出自exp/配置(README);本库不给训练循环,复现要配 Megatron-LM。
6. 代码地图(本章涉及)
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| grouped 前向 | megablocks/layers/dmoe.py | grouped_forward_once、grouped_permute_and_compute |
| 后端分发 | megablocks/layers/dmoe.py | ParallelDroplessMLP.forward_once |
| GroupedMLP | megablocks/layers/mlp.py | GroupedMLP.forward |
| 显存复用 | megablocks/layers/mlp.py | MemoryOptimizedGroupedMLP |
| GLU 变体 | megablocks/layers/glu.py | GroupedGLU、MemoryOptimizedGroupedGLU |
| 无补零重排 | megablocks/backend/kernels.py | gather、scatter |
| 依赖检测 | megablocks/grouped_gemm_util.py | assert_grouped_gemm_is_available |
| 注册表 | megablocks/layers/dmlp_registry.py | get、_REGISTRY |