18 · Switch Transformer: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity

MoE 稀疏计算 万亿参数 训练稳定性
Fedus, Zoph & Shazeer (Google) · 2021
arXiv:2101.03961

上一代的问题:Sparsely-Gated MoE (Shazeer et al., 2017) 使用 Top-2 路由,每个 token 激活多个专家,门控网络复杂;辅助损失设计繁琐,需要同时优化重要性损失和负载损失;训练极不稳定,需要对噪声和门控超参数精细调优;负载均衡策略不够可微,经常出现路由崩溃(所有 token 涌入同一个专家)。

这代改了什么:Top-1 路由 是核心简化——每个 token 只路由到一个专家,通信量减半,路由决策完全可微。辅助损失被统一为单个可微损失 loss = alpha * N * sum(f_i * P_i)。路由器用 float32 保精度,门控权重缩小初始化,专家 dropout 提至 0.4——这套稳定性配方的组合使得训练不再脆弱。

效果:第一个万亿参数 Transformer:Switch-C 1.6T。Switch-Base 训练速度是 T5-Base 的 7 倍,Switch-Large 比 T5-XXL 快 4 倍。多语言 101 种语言全部提升。知识蒸馏将 30% 的稀疏增益转移到稠密模型。PaLM 和 Gemini 的稀疏架构设计直接继承此工作。

一、核心简化:Top-1 Routing

Switch Transformer 最核心的设计决策是将 MoE 路由从 Top-2 简化为 Top-1:每个 token 只选择并路由到一个专家 FFN。这看似是退步——"每个 token 只有一个专家的意见,不是更不精确吗?"——但实践表明,不同 token 可以选择不同专家,模型整体从专家分化的组合中获取足够的表征多样性。

1.1 Top-1 vs Top-2 的数学对比

设门控网络输出分数向量 s = W_g · x,Softmax 后得到对各专家的概率分布 p = Softmax(s)

Top-2: y = pi · FFNi(x) + pj · FFNj(x)
Top-2 路由:加权组合两个专家的输出
Top-1: y = FFNargmax(p)(x)
Top-1 路由:只取分数最高的专家的输出

从计算角度:Top-2 需要 2x 的专家 FFN 前向计算 + 1x 的组合计算,且两个专家可能在不同设备上,需要额外的通信开销。Top-1 完全消除了组合步骤,通信量减半。

从表达能力角度:虽然单个 token 只经过一个专家,但不同 token 可以路由到不同专家。在一个包含 N 个专家、S 个 token 的 batch 中,模型实际上使用了 min(S, N) 个专家的组合来表达整个 batch 的分布。这种"每个 token 大致正确,整体通过组合取胜"的思想是关键。

1.2 路由崩溃风险的重新评估

Top-1 路由的一个直观风险是:如果路由决策错误,没有第二个专家来"兜底"。但实验发现,在适当的负载均衡约束下,路由崩溃不会发生——不同层倾向于产生不同的路由模式,底层层倾向于根据句法特征路由,高层根据语义特征路由。这种自发的专业化分工使得 Top-1 足够。

一个关键洞见:路由选择的"正确性"不需要每个 token 都完美——只要大部分 token 的路由是合理的,整体模型就能通过多层堆叠和 token 间的专家多样性获得足够的表达能力。追求每个 token 的"精确路由"反而会导致过度复杂化。

二、可微负载均衡损失

负载均衡是 MoE 的核心挑战:如果所有 token 都涌向同一个"赢家"专家,其他专家被闲置,模型的参数量优势就被完全浪费了。Switch Transformer 提出了 单个统一的辅助损失,替代了 Shazeer et al. 中繁琐的双目标损失。

2.1 损失公式

L = α · N · Σ fi · Pi
公式 1:Switch Transformer 负载均衡损失

其中:

2.2 为什么这个损失是可微的?

关键设计巧妙之处:fi 是离散的 token 计数(不可微),但 Pi 是通过 Softmax 输出的连续概率(完全可微)。两者的乘积 fi · Pi 在反向传播时,梯度可以流过 Pi 从而更新门控网络的权重。fi 作为一个加权系数,使得梯度信号与该 expert 的负载程度成正比——负载越重的专家,其对应的门控概率调整幅度越大。

乘以 N 是为了归一化:当 N 个专家均匀分配时,每个专家的 fi = 1/N,Pi = 1/N,乘积 N · fi · Pi = 1/N。当分布不均匀时,该值增大。最小化该损失等同于鼓励专家被均匀使用。

2.3 与 Shazeer et al. 损失的对比

特性Shazeer et al. (Top-2)Switch Transformer (Top-1)
损失项数2 项(重要性损失 + 负载损失)1 项(统一辅助损失)
噪声机制需要显式加噪(NormalNoise)来分散路由不需要——损失本身提供足够的梯度信号
可微性仅部分可微完全可微(通过 Pi 传递梯度)
超参数多个(noise_std, threshold, loss_weights 等)只有 alpha
训练稳定性敏感,需要精细调参稳定,默认 alpha=1e-2 即可工作

2.4 负载均衡的实际效果

实验中,使用该损失后各专家的 token 分配接近均匀(理想值是每个专家 1/N 的 token)。即使将 alpha 从 1e-2 向下调整到 1e-3,路由分布仍然相当均匀,说明该损失鲁棒性很好。论文还试验了 无辅助损失的极端情况:一些专家几乎不被使用,模型性能下降明显。

负载均衡与模型容量的关系:如果 64 个专家中只有 8 个被使用(路由崩溃),那么虽然模型参数量标注为"xxB",但有效容量只有 1/8。负载均衡损失正是为了防止这种"名义参数膨胀"。这是 MoE 论文中经常被忽视的关键指标。

三、训练稳定性技术

Switch Transformer 不仅简化了架构,还贡献了一套经过实战检验的 训练稳定性配方,后续几乎所有 MoE 模型都沿用了其中的部分或全部。

3.1 路由器精度:float32 保底

即使模型的其余部分使用 bfloat16 混合精度训练,路由器(门控网络)的权重和计算也必须保持 float32。原因很微妙:路由决策依赖于 Softmax 输出的概率排序,bfloat16 的低精度可能导致两个相近分数无法区分,从而改变路由选择。路由决策是离散的(选哪个专家),一个错误决策无法通过后续计算修正。

精度敏感性的根源:Softmax 的指数运算会放大数值差异。bfloat16 只有 ~7 位有效精度,当两个专家的门控分数差异在 10-4 量级时,bfloat16 无法区分,导致随机路由。float32 的 ~23 位有效精度完全够用。

3.2 缩小初始化

门控权重矩阵 Wg 使用比标准 Transformer 更小的初始化尺度。标准初始化通常采样自 N(0, 0.02),而 Switch Transformer 将门控初始化的标准差进一步缩小。目的是避免训练早期门控分数差异过大(导致路由选择过于确定),给所有专家一个被探索的机会。

这类似于强化学习中 exploration vs exploitation 的平衡:初期需要足够的随机性让专家都接收 token,后期再逐渐分化。

3.3 专家专用 Dropout

专家 FFN 层的 dropout 率设为 0.4,远高于非 MoE 层的典型值 0.1。原因在于:

在微调阶段,可将专家 dropout 降低或关闭以充分利用预训练学到的知识。

3.4 选择性精度训练

除了路由器的 float32 要求,Switch Transformer 还使用了 bfloat16 混合精度

历史语境:2021 年,bfloat16 刚在 TPU v2/v3 上广泛支持,混合精度训练还是一个相对前沿的话题。Switch Transformer 在这方面的经验为后续大型 MoE 模型(如 GLaM、PaLM)提供了直接的参考。

3.5 其他工程细节

技术描述目的
Expert cropping训练时随机丢弃部分专家的 token防止专家容量溢出并增强正则化
Capacity factor每个专家的 token 容量上限(通常设为 batch_size / N 的 1.0~1.25 倍)确保确定性内存分配和负载均衡
Layer normalization专家 FFN 前后均加 LayerNorm稳定深层梯度流动
Auxiliary loss deferral训练前几轮逐步增大 alpha让门控先学习基础路由,再施加均衡约束

四、分布式策略

Switch Transformer 的分布式训练策略是其实现万亿参数模型的关键。核心思路是将专家分布在多个设备上,结合数据并行和模型并行,形成 三维并行 方案。

4.1 专家并行 (Expert Parallelism)

与数据并行中每个设备保存完整模型副本不同,专家并行将不同的专家放置在 不同设备 上。设总共有 N 个专家、D 个设备,每个设备保存 N/D 个专家的参数。每个设备上的 standard layers(Self-Attention、Embedding)仍然是完整的,但专家 FFN 被分片。

一个 token 的典型路径:

  1. 本地完成 Self-Attention 计算
  2. 本地门控网络计算路由分数,确定目标专家所在设备
  3. 通过 All-to-All 通信 将 token 发送到目标设备
  4. 目标设备上的专家 FFN 处理 token
  5. 通过另一个 All-to-All 通信将结果返还给原设备
  6. 本地完成后续操作(Add & Norm、FFN 等)

4.2 All-to-All 通信模式

All-to-All 通信是整个分布式策略的瓶颈。与数据并行中的 All-Reduce(每个设备向所有设备广播)不同,All-to-All 是 置换通信:每个设备向每个其他设备发送不同的数据。

# 伪代码:Switch Transformer 的分布式前向传播
def distributed_switch_forward(x, experts, gate_network):
    # 1. 门控网络计算路由
    gate_scores = gate_network(x)           # [batch, N]
    expert_indices = argmax(gate_scores, dim=-1)  # [batch]

    # 2. 根据目标专家所在设备打包数据
    dispatch_data = prepare_all_to_all(x, expert_indices)

    # 3. All-to-All:将 token 发送到对应设备
    received_data = all_to_all(dispatch_data)

    # 4. 每个设备处理分配到的 token
    local_expert_output = local_experts(received_data)

    # 5. All-to-All:将结果返回原设备
    returned_data = all_to_all(local_expert_output)

    # 6. 根据原始位置重组输出
    output = combine(returned_data, expert_indices)
    return output

4.3 数据并行 + 模型并行 + 专家并行

在实际大规模训练中,三种并行策略叠加以充分利用硬件:

并行策略切分维度通信模式适用组件
数据并行batch 维度All-Reduce全部(梯度同步)
模型并行层/矩阵维度Peer-to-PeerSelf-Attention、FFN
专家并行专家维度All-to-All专家 FFN

三者的组合使得训练万亿参数模型成为可能:数据并行处理大的 batch 大小,模型并行解决单个大矩阵放不进一个设备的问题,专家并行(MoE 特有的)允许总参数量远超单个设备的存储能力。

4.4 Switch-C:1.6T 参数的具体配置

Switch-C 是论文中展示的最大模型,具体配置:

效率的关键:Switch-C 虽然总参数 1.6T,但推理时每个 token 只激活极少数参数(~1B),使得推理速度与同等参数量稠密模型相比快几个数量级。这就是条件计算的核心价值——用参数量换容量,但计算成本保持与激活参数成比例,而非总参数。

五、实验结果

5.1 训练速度对比

在相同 FLOPs 预算下,Switch Transformer 显著加速训练。核心实验结果以 T5-Base 为基线:

模型总参数量激活参数量速度提升 vs T5-Base速度提升 vs 同级 T5
T5-Base~220M~220M1x
Switch-Base~7B 激活 / ~157B 总~200M7xvs T5-Base 7x
T5-Large~770M~770M1x (vs T5-Large)
Switch-Large~26B 激活 / ~600B 总~700M4x+vs T5-Large 4x
T5-XXL~11B~11B1x (vs T5-XXL)
Switch-XXL~4xvs T5-XXL 4x
Switch-C~1.6T 总~1B~4xvs T5-XXL 4x

理解"速度提升":这里的 7x 并非 Switch-Base 比 T5-Base 快 7 倍(因为 Switch-Base 总参数大得多),而是——Switch-Base 达到 T5-Base 的精度所需的训练步数/时间减少了 7 倍。在相同的 FLOPs 预算下,Switch Transformer 每步的计算量比稠密 T5 更低(每个 token 只激活部分参数),所以每步更快。更大的加速比反映了稀疏架构更高的参数效率。

5.2 多语言性能

Switch Transformer 在 mT5(多语言 T5)的 101 种语言上进行了评估:

这意味着稀疏架构在处理 多模态和多语言 时尤其有价值——不同的专家可以隐式地专攻不同的语言特征、语法模式或主题领域。

5.3 知识蒸馏

论文还探索了将稀疏 Switch Transformer 的知识蒸馏到稠密模型:

蒸馏的深层含义:如果 30% 的稀疏增益可以转移到稠密模型,那推理时就不需要维护 MoE 架构了——部署稠密蒸馏模型更简单。但如果追求极致性能,直接使用稀疏 Switch Transformer 推理仍然是必要的。这种 trade-off 在实际工程中经常权衡。

5.4 消融实验:Top-1 vs Top-2

论文进行了详细的消融实验来验证 Top-1 的有效性:

配置路由策略速度质量 (GLUE 平均)
Switch-Base (Top-1)Top-1最快83.1
MoE-Base (Top-2)Top-2较慢(2x 专家计算)83.0
Dense-Base无 MoE与 Top-1 接近81.4

Top-1 不仅比 Top-2 更快(通信减半、无组合计算),而且质量相当。这有力地证明了"简单就是最好"——复杂的 Top-2 路由并不比简单的 Top-1 路由带来额外的好处。

六、上下游关联

6.1 上一篇:Sparsely-Gated MoE (Shazeer et al., 2017)

Shazeer 的开创性工作首次将 MoE 条件计算引入深度学习,证明了 137B 参数 LSTM MoE 的可行性。但它的 Noisy Top-K 门控 + 双辅助损失 方案过于复杂,训练不稳定,难以推广到 Transformer。Switch Transformer 站在其肩膀上,做了关键的"简化"工作。

6.2 下一篇:DeepSeek MoE (2024)

DeepSeek 将 MoE 进一步演化:引入 细粒度专家(细分成更小的专家)、共享专家(所有 token 必定访问的公共知识)和 辅助损失无关的负载均衡(动态调整偏置代替辅助损失)。Switch 的 Top-1 简化哲学在 DeepSeek 中得到了继承和扩展。

6.3 采用此工作的模型

6.4 影响总结

Switch Transformer 是 MoE Transformer 的 "定型之作"。在此之前,MoE 是一个有趣但难以实用的小众技术;在此之后,MoE 成为训练大模型的主流方案之一。2024 年发布的多个最强开源模型(DeepSeek V2/V3、Mixtral、Qwen-MoE)均采用 MoE 架构,全部从 Switch Transformer 获得了核心设计元素。

七、个人思考

7.1 "简单性"本身就是创新

Switch Transformer 最令人印象深刻的不是发明了什么新技术,而是 敢把一个复杂的系统简化到极致。Top-1 路由这个想法简单到几乎"愚蠢",但 Shazeer 等人之前没有人敢这么做。这提醒我们:在深度学习研究中,"简化"往往比"复杂化"更难,但也更有价值。

7.2 训练稳定性配方的积累价值

论文中关于 training stability 的细节(float32 router、小初始化、专家 dropout)单独看都像是"小技巧",但组合起来构成了 MoE 模型的 最佳实践配方。这种配方的价值不低于架构创新本身——没有稳定性配方的架构是"空中楼阁"。

7.3 从"参数膨胀"到"参数效率"的范式转移

Switch Transformer 标志着对模型规模思考的转变:不再只是追求更大的参数数量,而是追求 更高的参数效率——如何让更多的参数但更少的计算达到更好的效果。Switch-C 的 1.6T 参数只有 ~0.1% 被每个 token 使用,有效参数量只是个位数 B。这表明未来的方向是"稀疏激活 + 稠密质量"。

7.4 简化的代价:专家专业化的局限性

Top-1 的代价是每个 token 失去了获得多方面知识的能力。虽然模型整体可以补偿,但某些需要跨领域知识的任务(如科学推理)可能仍需要更复杂的路由。这也解释了为什么后续模型(如 Mixtral 8x7B)回到了 Top-2——不同的路由策略适用于不同的规模和任务类型。

7.5 论文写作的启示

Switch Transformer 论文的写作值得学习:它把核心贡献放在最前面(Figure 1 就是 Switch 架构对比),实验设计清晰(每个实验只回答一个明确的问题),消融完备。相比许多晦涩的论文,这种"结果导向"的写作风格使得它的影响力远超其技术创新的"重量"。