跳到主要内容

数据截至 (上游 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_oncemegablocks/layers/dmoe.py:239-258)。


4. 真实实现

4.1 前向:元数据都回到「无补零」版本

grouped_forward_oncemegablocks/layers/dmoe.py:239-258)调用的是第 1 章的 indices_and_bins(不算 padded_bins),然后 grouped_permute_and_computemegablocks/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.forwardmegablocks/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)。

gggrouped_gemm_utilmegablocks/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.0setup.py:72)。

4.3 手写 autograd:激活重算 + 显存复用

朴素版 forward 要存两个中间量(gmm 输出、激活输出)。MemoryOptimizedGroupedMLPmegablocks/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-480megablocks/layers/gelu.py

MemoryOptimizedMLP(sparse 版,megablocks/layers/mlp.py:187-306)是同一思路在 stk 算子上的版本;GLU 也有对应的 MemoryOptimizedGroupedGLUmegablocks/layers/glu.py:71-170),但 sparse GLU 的 memory_optimized 直接 NotImplementedErrormegablocks/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.0grouped_gemm 0.3.0
拓扑计算每前向建稀疏拓扑 + 转置不需要
host 同步padded_bins[-1].cpu()batch_sizes.cpu()
官方定位老路径默认 + Hopper 推荐(README)

默认配置已经回答了这个问题:mlp_impl: str = 'grouped'megablocks/layers/arguments.py:50)。注册表在 dmlp_registrymegablocks/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 传成 binsmegablocks/backend/kernels.py:206-208scatter),所以第 2 章的索引数学这里完全适用。
  • 显存优化版是训练默认假设。 测试里 memory_optimized_mlp=Truetests/layers/dmoe_test.py:73);关掉它就走朴素 gmm 三行。
  • 性能数字看 benchmark。 仓库自带 megablocks/ops/matmul_benchmark.pypermute_benchmark.py 等微基准,论文的 40% / 2.4x 数字出自 exp/ 配置(README);本库不给训练循环,复现要配 Megatron-LM。

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

主题文件路径符号名
grouped 前向megablocks/layers/dmoe.pygrouped_forward_oncegrouped_permute_and_compute
后端分发megablocks/layers/dmoe.pyParallelDroplessMLP.forward_once
GroupedMLPmegablocks/layers/mlp.pyGroupedMLP.forward
显存复用megablocks/layers/mlp.pyMemoryOptimizedGroupedMLP
GLU 变体megablocks/layers/glu.pyGroupedGLUMemoryOptimizedGroupedGLU
无补零重排megablocks/backend/kernels.pygatherscatter
依赖检测megablocks/grouped_gemm_util.pyassert_grouped_gemm_is_available
注册表megablocks/layers/dmlp_registry.pyget_REGISTRY