跳到主要内容

数据截至 (上游 commit 952db33d6eac)

MegaBlocks — 架构与原理

30 秒导读: MegaBlocks 是 Databricks 开源的 MoE(混合专家)训练计算库。它只做一件事并做到极致:把 MoE 层里「路由完每个专家拿到的 token 数不一样」这个对 GPU 极不友好的计算,重构成 block-sparse / grouped GEMM 能高效执行的形态——不用像传统 MoE 那样丢掉溢出的 token(dropless),速度反而比丢 token 的实现快最多 40%。它是「层实现」,不是训练框架,生产上嵌在 Megatron-LM 里用。


1. 这是什么(零基础也能懂)

一句话定义

MegaBlocks 是一个 PyTorch 库,提供两种 MoE 层实现——标准的 MoE(可能丢 token)和它主打的不丢 token 的 dMoE(dropless MoE)——以及让这两种层跑快所需的全部底层算子(排序、直方图、前缀和、gather/scatter、block-sparse 矩阵乘)。

它要解决谁的什么问题

假设你在训一个大语言模型,想把 FFN 层换成 MoE 层来扩大参数量:

  • 你有一个「路由器」,给每个 token 打分,把它送给 top-k 个专家(每个专家是一个独立 FFN)。
  • 每一步里,每个专家分到多少 token 是随数据变化的——有的专家热门、有的冷门。
  • GPU 的矩阵乘(GEMM)喜欢大而规整的输入;「8 个专家、每个专家的 token 数分别是 901 / 37 / 412 / …」这种变长输入,直接算就是一堆碎小 GEMM,慢得没法用。

业界的传统解法是给每个专家定死容量(capacity),多出来的 token 直接丢掉不算——这就是 token dropping。它能跑快,但丢 token 伤效果,还多出一个难调的 capacity_factor 超参。

MegaBlocks 要的就是「既不丢 token,又不牺牲硬件效率」。 它的答案是把变长问题重构成 block-sparse 矩阵乘(详见第 2 章)。

它能做什么

能力具体形态位置
标准 MoE 层MoE / ParallelMLP,定容量桶,可能丢 tokenmegablocks/layers/moe.py:440
dropless MoE 层dMoE / ParallelDroplessMLP,不丢 tokenmegablocks/layers/dmoe.py:323
两种专家计算后端block-sparse(stk)/ grouped GEMMmegablocks/layers/mlp.py:308:499
路由与负载均衡LearnedRouter + load balancing loss + z-lossmegablocks/layers/router.py:61
专家并行数据/专家/流水线并行,配 Megatron-LMmegablocks/layers/moe.py:237
底层 CUDA/Triton 算子sort / histogram / cumsum / gather / scattermegablocks/ops/csrc/

用起来什么样

最小用法(单机、8 个专家、top-2),取自测试的构造方式(tests/layers/dmoe_test.py:47-88):

import torch
from megablocks.layers.arguments import Arguments
from megablocks.layers.dmoe import dMoE

args = Arguments(
hidden_size=512,
ffn_hidden_size=1024,
moe_num_experts=8,
moe_top_k=2,
mlp_impl='grouped', # 推荐路径;需 pip install megablocks[gg]
bf16=True,
)
layer = dMoE(args).cuda().to(torch.bfloat16)

x = torch.randn(1024, 16, 512, device='cuda', dtype=torch.bfloat16) # [sl, bs, hs]
out, bias = layer(x) # 默认 return_bias=True,返回 (输出, bias)

注意三件事:

  • 输入布局是 [sequence, batch, hidden](Megatron 习惯),不是 [batch, sequence, hidden]
  • dMoEMoE 的对外接口完全一样,dMoE 只是把内部的专家计算换成 ParallelDroplessMLPmegablocks/layers/dmoe.py:323-326)。
  • 在 Megatron-LM 里训练时,用 from_megatronmegablocks/layers/arguments.py:95-99)把 Megatron 的 args 一键转成 Arguments,启动脚本见 exp/dmoe/dmoe_125m_8gpu.sh

一句话直觉

把 MoE 层想成一个分拣中心。 每个 token 是一个包裹,路由器按地址(专家 id)分拣。传统做法是给每个分拣口一个固定大小的筐:筐满拒收(丢 token)、筐空也得占地方。MegaBlocks 的做法是:每个分拣口的队列向上取整到 128 个一摞——多垫一点空包裹,但换来所有队列都能用标准化的「板条箱流水线」(block-sparse GEMM)处理,一个包裹都不丢。


2. 顶层全景(它大概怎么转)

2.1 一张图看懂 dMoE 前向

怎么读这张图: 从左到右是数据流。左边是路由,中间是「把变长变规整」的核心三步,右边是还原。①~④ 是每一步都要做的元数据计算,全部在 torch.no_grad() 里。

输入 x [sl, bs, hs]


① LearnedRouter Linear(hs → E) → softmax → top-k
│ 产出: scores, expert_weights, top_experts

② 分拣元数据(no_grad) sort → histogram → cumsum → round_up(128) → cumsum
│ 产出: indices, bin_ids, bins, padded_bins

③ padded_gather 按排序结果把 token 重排,每个专家补零到 128 的倍数
│ x: [tokens, hs] → [padded_tokens, hs]

④ 专家 MLP(两种后端之一)
├─ sparse: stk sdd → act → stk dsd (block-sparse 矩阵乘)
└─ grouped: gmm → act → gmm (grouped GEMM,免补零)


⑤ padded_scatter 乘路由权重、按原顺序散射回去、top-k 求和


输出 [sl, bs, hs](+ 可选 bias、shared expert 输出)

2.2 部件职责

部件干什么在哪个文件
MoE / dMoE层的对外门面:autocast → 路由 → 专家计算 → shared expert 融合megablocks/layers/moe.py:440megablocks/layers/dmoe.py:323
LearnedRouter一个无 bias 的 Linear,产出每个 token 的专家分数与 top-k 分配megablocks/layers/router.py:61
ParallelMLP标准 MoE 的专家计算:定容量桶 + torch.bmm,含专家并行通信megablocks/layers/moe.py:96
ParallelDroplessMLPdMoE 的专家计算:补零到 128 的倍数 + block-sparse / groupedmegablocks/layers/dmoe.py:18
MLP / SparseMLP / GroupedMLP三种专家 MLP 实现(bmm / stk / grouped_gemm)megablocks/layers/mlp.py:91:308:499
ops.sort / histogram / inclusive_cumsum分拣元数据三件套(cub CUDA kernel)megablocks/ops/sort.py:26histogram.py:20cumsum.py
ops.padded_gather / padded_scatter带补零的重排/还原(Triton kernel)megablocks/backend/kernels.py:107:173
ops.binned_gather / binned_scatter定容量桶的重排/还原(标准 MoE 用)megablocks/backend/kernels.py:392:421
ops.topology生成 block-sparse 矩阵的列索引megablocks/ops/topology.py:20csrc/indices.h
all_to_all专家并行的跨卡 token 交换(带梯度的 autograd 封装)megablocks/layers/all_to_all.py:8
mpu专家并行拓扑:专家维 × 隐藏维的二维切分megablocks/layers/mpu.py:65-93
Arguments全部配置(模型/MoE/并行/计算),from_megatron 适配megablocks/layers/arguments.py:23

2.3 主线走一遍(一次 dMoE 前向)

dMoE.forwardmegablocks/layers/moe.py:459-475)为线,高层不进代码:

  1. 降精度先行。 cast_if_autocast_enabled 先把激活转成 bf16/fp16,这样后面的 token 重排搬的是半精度数据,省带宽(megablocks/layers/moe.py:463-464 的注释明说这一点)。
  2. 路由。 LearnedRouter 算出 scoresexpert_weightstop_experts 三件套(megablocks/layers/router.py:92-113)。
  3. 分拣元数据。top_experts 做 radix sort 得到重排索引 indices,histogram 得每专家 token 数,cumsum 得桶边界 bins;再向上取整到 128 的倍数得到 padded_binsmegablocks/layers/dmoe.py:131-154)。
  4. 重排 + 专家计算 + 还原。 padded_gather 把 token 按专家聚拢并补零 → 专家 MLP(stk 或 grouped GEMM)→ padded_scatter 加权散射回原位置、top-k 维求和(megablocks/layers/dmoe.py:156-193)。
  5. 可选附加。 训练时把 (tokens_per_expert, scores) 存进全局列表供负载均衡损失用(megablocks/layers/moe.py:430-431);开了 shared_expert 再叠一个所有 token 共享的稠密 MLP(megablocks/layers/moe.py:468-473)。

开了专家并行(moe_expert_model_parallelism=True)时,第 4 步换成 parallel_forward_once:先本地重排、再 all_to_all 跨卡交换、再本地重排,算完原路返回(megablocks/layers/moe.py:237-423),详见第 4 章。


3. 阅读地图(建议顺序)

四章由浅入深。赶时间读 01 → 02 就能抓住 MegaBlocks 的全部要害;要上手用再读 03,要多卡训练再读 04。

顺序章节讲什么适合谁
101-why-moe-is-hard.mdMoE 为什么难算:变长专家输入、capacity factor、token dropping、标准 MoE 的 binned 实现所有人必读,这是理解 dMoE 的前提
202-dropless-block-sparse.mddMoE 的 block-sparse 重构:补零到 128、stk 稀疏拓扑、sdd/dsd想搞懂论文核心想法怎么落地的人
303-grouped-gemm.mdgrouped GEMM 路径:免补零的变长分组乘、显存复用要上手训练、关心 Hopper 性能的人
404-router-and-parallelism.md路由细节、负载均衡损失、专家并行的三次置换 + all_to_all、Megatron 集成要多卡训练、调收敛的人

4. 巧妙之处(可借鉴的技术)

每条先白话点出妙处,细节见对应章节。

  • radix sort 只排需要的位。 排序 key 是专家 id,最多 ceil(log2(num_experts)) 位,于是告诉 cub radix sort 只排这些位(sort_end_bitmegablocks/layers/moe.py:110)。64 个专家只排 6 位,比对 32 位整数全排快得多。见第 1 章。
  • 稀疏矩阵的「行偏移」是免费的等差数列。 补零后每个 128-token 行块恰好属于一个专家,所以每行的非零块数是常数,offsets 不用算,直接 torch.arange 生成(megablocks/layers/dmoe.py:82-89)。真正要算的只有列索引。见第 2 章。
  • 转置拓扑用「按列排序行索引」推出来。 backward 需要转置后的稀疏拓扑,做法是拿 column_indices 再排一次 radix sort,用排序的 gather 顺序重排 row_indices(megablocks/layers/dmoe.py:35-66)。见第 2 章。
  • gather 和 scatter 是同一个 Triton kernel 的两个方向。 _padded_copyA_TO_B 编译期标志切换方向,前向 gather 的 backward 就是 scatter(不带权重),反之亦然(megablocks/backend/kernels.py:45-104megablocks/ops/padded_gather.py:29-42)。见第 2 章。
  • 手写 autograd 做显存复用。 MemoryOptimizedGroupedMLP 在 backward 里重算激活函数(rematerialize)、复用同一块显存放中间梯度,省掉一整个中间激活的显存(megablocks/layers/mlp.py:397-496)。见第 3 章。
  • 初始化与并行方式无关。 每个 rank 先建完整主权重、用同一种子初始化、再切出自己那片,保证「数据并行」和「专家并行」同一 seed 得到同一组权重(megablocks/layers/mlp.py:43-88)。见第 4 章。

5. 边界与局限

诚实地说,这个库的适用范围相当窄,而且有些坑是代码里写死的:

  • 它是层实现,不是训练框架。 训练循环、数据加载、优化器全靠宿主(Megatron-LM 或 LLM-Foundry);负载均衡损失只「存进全局列表」(megablocks/layers/moe.py:17-29),怎么加进总 loss 是宿主的事。
  • 维度被 128 绑死。 block-sparse 路径要求 ffn_hidden_size 能被 128 整除(self.blocking = 128megablocks/layers/dmoe.py:24),否则 topology 直接 raise(megablocks/layers/dmoe.py:71-75)。
  • sparse 后端与新 Triton 不兼容。 mlp_impl='sparse' 在 triton >= 3.2.0 下直接报错,官方让走 groupedmegablocks/layers/arguments.py:75-85);grouped 也是默认(megablocks/layers/arguments.py:50)。
  • 额外依赖。 sparse 路径依赖 stanford-stk==0.7.1setup.py:66),grouped 路径依赖 grouped_gemm==0.3.0setup.py:72),两者都是同一作者的小众库;CUDA 扩展只编译了 sort/histogram/cumsum/indices/replicate 几个小 kernel(csrc/ops.cu:10-19)。
  • 路由器本身不并行。 注释明说每个设备都持有完整路由权重,因为每个设备都要给自己的 batch 路由(megablocks/layers/router.py:66-69)。专家数极大时这是一份每卡都有的冗余。
  • grouped 路径有 host 同步点。 GroupedMLP.forward 要把 tokens_per_expert 搬到 CPU(megablocks/layers/mlp.py:502),并行路径里为此特意复用已搬好的 CPU 张量(megablocks/layers/moe.py:383-390)。见第 3 章。
  • 低提交频率。 仓库健康度一般,README 自称 "light-weight";把它当「设计参考 + 可移植算子」比当「长期维护的框架」更合适(这是书架侧判断,非代码事实)。

6. 横向对比

同书架上与 MegaBlocks 相邻的项目,取舍各不相同:

项目层级与 MegaBlocks 的关系
Megatron-LM训练框架MegaBlocks 的宿主:它提供 MoE/dMoE 层,Megatron 提供数据/流水线并行、训练循环与启动脚本(exp/
verlRL 训练框架同样站在 Megatron 系生态上做训练,但解的是 RL 的训推协同;MoE 模型经 Megatron 后端时可受益于这类层实现

与书架外的同类比:Tutel(微软)是 token-dropping MoE 的代表实现,README 报告 dMoE 比 Tutel 最优 capacity_factor 配置快最多 40%(README.md);Switch/GShard 是 capacity factor 范式的提出者,第 1 章的标准 MoE 实现就是这套范式的落地。MegaBlocks 的独特站位是「用 block-sparse 把 dropless 做到不亏性能」,这个思路后来被多个 MoE 训练栈吸收。


7. 代码地图(导航索引)

读源码的建议入口顺序:layers/moe.pylayers/dmoe.pybackend/kernels.pylayers/mlp.py。各章末尾有更细的地图。

主题文件路径符号名
层门面(MoE/dMoE)megablocks/layers/moe.pyMoEParallelMLP.forward
dropless 层megablocks/layers/dmoe.pydMoEParallelDroplessMLP
路由器megablocks/layers/router.pyLearnedRouter.forwardbatched_router_zloss
分拣元数据megablocks/layers/moe.pyParallelMLP.indices_and_bins
补零元数据megablocks/layers/dmoe.pyParallelDroplessMLP.indices_and_padded_bins
稀疏拓扑megablocks/layers/dmoe.pytopologysparse_transpose
gather/scatter Triton kernelmegablocks/backend/kernels.py_padded_copypadded_gatherpadded_scatter
定容量桶 kernelmegablocks/backend/kernels.py_binned_copybinned_gatherbinned_scatter
专家 MLP 三实现megablocks/layers/mlp.pyMLPSparseMLPGroupedMLP
显存复用 autogradmegablocks/layers/mlp.pyMemoryOptimizedMLPMemoryOptimizedGroupedMLP
专家并行megablocks/layers/moe.pyParallelMLP.parallel_forward_once
并行拓扑megablocks/layers/mpu.pyexpert_sharding_degreehidden_sharding_degree
CUDA 小 kernelcsrc/sort.hcsrc/histogram.hcsrc/indices.hcub_radix_sortcub_histogramConstructIndicesKernel
配置megablocks/layers/arguments.pyArgumentsfrom_megatron