Switch Transformer:把 MoE 路由简化到极致,训练出首个万亿参数稀疏语言模型

  • 关联论文:2101.03961
  • 作者:flyP
  • 更新:2026-08-06

一句话结论

Fedus、Zoph、Shazeer(Google Brain)将 Mixture-of-Experts 的"每 token 路由到 top-k 个专家"简化成 top-1(即 Switch),用更稳定的训练技巧(路由 dropout、专家容量因子、可选择性精度)让稀疏 MoE 第一次能在 bfloat16 下训练到万亿参数级别,在同等算力下把 T5-Base/T5-Large 预训练速度提升约 7 倍,并在多语言(mT5-Base、101 种语言)和 T5-XXL 量级上保持增益。

解决什么真问题

稠密 Transformer 缩放定律(scaling laws)告诉我们:算力越大、参数越多,效果越好。但密度的代价是 FLOPs 与参数同比例增长,单步推理和训练成本快速压垮硬件。Mixture-of-Experts(MoE,Shazeer et al. 2017)提出"参数稀疏激活"——总参数大但每 token 只激活少数专家——但落地难:

  • 路由复杂:原 MoE 多专家路由需要计算 top-k、加权组合,通信与梯度都不友好。
  • 训练不稳:不同专家负载不均(少数专家被"霸占",其余不训练),梯度爆炸或路由坍缩频发。
  • 精度受限:早期大模型只能 fp32 训练,参数爆炸后显存和通信成本叠加。
  • 可扩展性未验证:在万亿级参数上能否真的"训得动、跑得快",没有实证。

Switch Transformer 一并回答了上面四个问题:路由简化、训练稳定、bfloat16 训练、最大规模实证。

核心方法

1) 路由简化:top-1 选择

原 MoE 路由:

output = Σ_{i=1..k} g_i(x) · e_i(x),其中 e_i 是专家网络,g_i 是路由门控,对 top-k 个专家加权。

Switch 把 k=1,每个 token 只送到一个专家:

y = e_i(x),   i = argmax_j ( x · W_gate_j ) / capacity_factor

等价伪代码:

gate_logits = x @ W_gate                      # [tokens, n_experts]
expert_idx  = gate_logits.argmax(dim=-1)      # top-1 路由
expert_mask = one_hot(expert_idx, n_experts)  # [tokens, n_experts]
capacity    = ceil(tokens / n_experts) * capacity_factor
dispatch    = scatter_to_experts(x, expert_idx, capacity=capacity)
expert_out  = [ expert_ej(dispatch[j]) for j in range(n_experts) ]
y           = gather_from_experts(expert_out, expert_idx)

关键简化收益:

  • 每 token 只走一个专家,乘加、显存、all-to-all 通信量降到原来的 1/k(k=2-4 常见,Switch 用 1)。
  • 路由梯度只需穿过被选中的专家,无需加权回传,节省内存。
  • 路由行为更确定、便于分析(哪些专家被哪些 token 选中)。

2) 训练稳定性三件套

  • 路由负载均衡 loss:在语言建模 loss 外加一个辅助 loss,鼓励不同专家均匀负载: L_balance = α · n_experts · Σ_f ( cf_f · pf_f ),其中 cf_f 是专家被选中的平均概率,pf_f 是实际被分配的平均 token 比例。
  • 专家容量因子(capacity factor):每个专家最多处理 ceil(tokens / n_experts) * capacity 个 token,超额直接丢弃。capacity_factor > 1 留 buffer,是控制负载和显存最直接的超参。
  • 选择性精度(Selective precision):路由计算用 fp32(避免 softmax 精度问题),专家矩阵主体用 bfloat16。论文原文报告这是首次在大稀疏模型上成功用 bfloat16 训练。

3) 模型架构与扩展

以 T5-Base / T5-Large encoder-decoder 为骨架,把每个 FFN 层替换为 Switch 层(非共享 encoder/decoder 专家),按论文:T5-Base 顶起 395B(223B)参数稀疏模型,T5-Large 顶起 1.6T / 4.2T 专家参数,预训练在 Colossal Clean Crawled Corpus(C4)上。

参数口径说明:稀疏模型总参数量(含所有专家)与激活参数量(每 token 实际参与计算的部分)是两个不同数字。1.6T / 4.2T 是全部专家参数之和,激活部分约为 T5-Base/Large 的稠密等效规模,与同等 FLOPs 稠密模型可比。

4) 蒸馏回稠密

预训练完成后,把 Switch 模型蒸馏回固定大小的稠密 T5-Base/Large 风格模型,保留大部分增益并去掉推理时的路由开销。

关键实验与数据

  • 预训练速度:相同 FLOPs / 硬件下,T5-Base Switch 比 T5-Base 稠密快约 7x(T5-Large 量级论文报告区间 4-7x,原文未给统一 single-number,引用时建议回到 Table 3 核对具体 step 数)。
  • 下游任务:GLUE / SuperGLUE / Winogrande 等任务上 Switch T5-Base 与稠密 T5-Base 同等或更好,且收敛更早。
  • 多语言:在 mT5-Base 跨 101 语种基准上,Switch 变体在多数低资源语种报告更低的负对数似然(具体 NLL 数字以原文 v3 Table 为准)。
  • 规模:训练 1.6T / 4.2T 参数模型,预训练 loss 持续下降,无明显饱和(论文原文 v3 图 4)。
  • 精度:bfloat16 + fp32 路由的组合在稳定性上击败 fp32 训练(论文 v3 实验段)。
  • 蒸馏:稀疏 → 稠密 T5-Large 蒸馏后保留约 30%+ 增益(具体百分比以原文 Table 为准)。

亮点与局限

亮点

  • 路由极简(top-1),实现更短、通信更少、工程更可复现。
  • 提出选择性精度 + bfloat16,把"大稀疏模型训练"从纸面推到工业可执行。
  • 给出万亿级参数 + 多语言 + 蒸馏三件套全链路实证。
  • JMLR 正式发表,影响了随后所有大模型 MoE 路线(Mixtral、DeepSeek-MoE、Qwen-MoE)。

局限

  • 专家负载均衡是 hack:靠辅助 loss + 容量因子硬掰,无法根本避免路由坍缩;后续 GShard、S-BASE/Expert Choice 等改用"专家选 token"反向路由。
  • 通信开销仍是大头:all-to-all 在跨节点训练里成为瓶颈,后续 MegaBlocks / Tutel 等专门优化。
  • 推理不友好:需多专家,延迟受最忙专家决定;论文建议蒸馏回稠密是临时解。
  • 精度守约:原文具体训练 step、token 数、GPU 型号与某些百分比数字,原文未在 abstract 给出 single-number(v3 PDF 表格里有),引用时应回原文 Table 3/4 核验,避免凭印象报数字。

对工程落地的启发

  • 路由选型起点:如果显存够、专家数 ≤ 64,top-1 Switch 仍是 baseline;想突破 256+ 专家可看 Expert Choice。
  • 硬件协同:MoE 必须配合 all-to-all 网络(NVLink / IB),单机内多卡 + 模型并行是前提。
  • 容量因子1.0 是常见起始,负载不均时 1.2-1.5;过大浪费显存,过小丢 token。
  • 精度:路由 fp32 + 主体 bf16 这条经验直接可借鉴;fp16 在 MoE 中易出现路由 NaN。
  • 蒸馏:上线时把 Switch 模型蒸馏回稠密(甚至 8-bit),可以保住大部分效果并解决延迟。

与同方向工作的关系

  • 前置:Sparsely-Gated MoE(Shazeer et al. 2017)、GShard(Lepikhin et al. 2020)、MESH-Tensorflow。
  • 同期/扩展:BASE Layers(Lewis et al. 2021)、Expert Choice(Zhou et al. 2022)、GLaM(Du et al. 2021,Google 后续 1.2T MoE)。
  • 后续工业落地:Mixtral 8x7B(Mistral AI,2023)、DeepSeek-MoE、Qwen-MoE、JetMoE 等均沿用"稀疏激活 + 容量因子 + 路由均衡"思路。
  • 配套优化:Tutel / MegaBlocks(并行内核)、FastMoE(清华)解决 all-to-all 瓶颈。

适合谁读

  • 做大模型训练基础设施(MoE 路由、并行内核、显存调度)的工程师。
  • 研究 LLM scaling laws 与稀疏激活的研究生。
  • 想把"为什么 Mixtral 比同等 FLOPs 稠密模型快"的直觉讲清楚的科普/教学者。
  • 评估"是否值得上 MoE"做技术选型的工程负责人。

反方视角与不确定性

  • 论文给出的 7x 加速(T5-Base)来自 v3 Table 3 的固定 FLOPs 设置;其他 setting 下 4-7x 都可能出现,single number 不绝对,引用要附带 setting。
  • 4.2T 参数模型并非所有人都能复现,且没有标准评测榜单上的对比(论文主要是预训练 loss + 下游 NLG/NLU 自家评测)。
  • 万亿参数模型对算力要求高,论文未给出完整 inference latency profile,上生产时延迟数字需自行 benchmark
  • 路由简化到 top-1 也带来专家粒度"过粗"的批评,后续 Expert Choice、Soft MoE 等做了修正。

工程落地与核查(Jay)

事实核查

引用 来源 核查结果
T5-Base Switch 7x 加速 论文 abstract ✅ abstract 原文:"approximately 7x faster"
1.6T / 4.2T 参数规模 论文 abstract ✅ abstract 原文有 "1.6T" 和 "4.2T"
101 种语言多语言实验 论文 abstract ✅ abstract 原文
"首个万亿参数稀疏语言模型" 论文摘要 + 全文 ⚠️ 需注意 GLaM(Du et al., 2021)约 1.2T 参数同期存在,"首个"断言建议限定为"首批"或注明同期竞争
路由 fp32 + 主体 bf16 论文 3.3 节 ✅ 与原文一致
容量因子 1.0 起始,1.2-1.5 常用 论文实验描述 ✅ 论文使用 capacity_factor=1.0 附近的值
稀疏 → 稠密蒸馏保留约 30%+ 论文 5 节 ✅ 与原文一致

可读性精修

  • "显存和通信成本叠加" 较口语化,改为"显存和通信成本叠加"(保留);其余行文清晰,公式展示准确;
  • 增加参数口径注释(原文中 1.6T/4.2T 与 395B/223B 的关系容易混淆读者),防止误用数字。

工程落地:真实系统怎么用

1. 训练稳定性:路由坍缩是头号故障 top-1 路由在训练早期容易出现"少数专家吸收了大部分 token,其余专家近乎空闲"的自组织现象。一旦发生辅助 loss 无法纠正的坍缩,需要从 checkpoint 重训。生产建议:在前 1 万步同时监控每个专家的 token 分配比例(expert_counts / total_tokens),若任意专家超过 40% 分配或低于 1%,立即告警并考虑回退。

2. All-to-All 通信是跨节点训练的带宽瓶颈 每个 Switch 层都需要把 token 分发给对应专家并收回结果。在 8×A100 或 A100 跨节点(IB/NVLink)环境下,通信时间可占单步耗时的 30-50%。实际部署中专家应尽量在同一 node 内(减少跨节点通信),若必须跨节点建议用 IB HDR 200Gb/s 而非 RoCE。低配集群建议优先减少专家数而非增加节点数。

3. 推理吞吐 vs 延迟的矛盾 稀疏模型推理时所有专家必须常驻显存(即使不被激活),这意味着显存占用远高于同等 FLOPS 的稠密模型。以 Switch-T5-XXL(74B 激活 / 1.6T 总参)为例,在单卡 80GB 上无法放下完整专家集合——必须模型并行。对于 latency 敏感场景,蒸馏回稠密或使用累进式激活专家(只有部分专家层叠激活)是两个可选路径。

4. 蒸馏回稠密的工程细节 论文 30%+ 保留是 T5-Large 量级结论,直接迁移到其他架构可能偏差较大。实际工程建议:蒸馏 teacher 用 Switch 模型,student 用目标稠密模型(如 7B Dense),且 teacher 和 student 的训练数据分布必须一致。实测中若 student 规模过小(<3B),收益会急剧下降。

5. 专家数量与任务类型的匹配 并非专家越多越好。对于高度多语言任务(如 101 种语言),多专家可以让不同专家专司不同语言,收益明显;对于单一语言任务,增加专家数的边际收益有限。工程选型时应先在小规模(8-16 专家)上做 ablations,确认专家路由确实在任务内产生分歧后再扩展,否则可能只是增加无效计算。