首页
/ Switch Transformers 稀疏 MoE 模型在 Transformers 中的实现解析:路由机制、专家调度、负载均衡损失与 8-bit 量化实战

Switch Transformers 稀疏 MoE 模型在 Transformers 中的实现解析:路由机制、专家调度、负载均衡损失与 8-bit 量化实战

2026-09-07 17:29:54作者:尤辰城Agatha

Switch Transformers 是一种把 T5 的稠密 MLP 层替换为 Mixture-of-Experts(MoE)的稀疏编码器-解码器模型,Transformers 仓库完整实现了其 Top-1 路由、专家容量控制与负载均衡损失等核心机制。本文以仓库官方文档 docs/source/en/model_doc/switch_transformers.md 为主线,结合 模型实现源码配置类源码,系统讲解从推理调用、配置参数到路由与损失函数的底层细节,并给出 bitsandbytes 8-bit 量化的可运行示例。读完后你可以直接加载 google/switch-base-8 等原始检查点进行生成,也能理解每个稀疏化配置项在源码中的确切作用。

1. 什么是 Switch Transformers

按官方文档的定义,Switch Transformers 是一个稀疏 T5 模型:其 MLP 层被替换为 Mixture-of-Experts 结构,一个路由机制(routing mechanism)将每个 token 关联到一个专家(expert),而每个专家就是一个稠密 MLP。稀疏性(sparsity)带来更好的扩展能力,而路由机制允许模型在推理时按需选择相关权重,从而在不增加单 token 计算量的前提下大幅提升模型容量。

文档同时指出,所有官方原始检查点都收录在 Hugging Face 的 Switch Transformer 集合中(如 google/switch-base-8),模型由 ybelkada 与 ArthurZ 于 2022-11-15 贡献进 Transformers。

从源码结构看,该实现包含以下核心组件(均位于 src/transformers/models/switch_transformers/ 目录):

组件 文件位置 职责
SwitchTransformersConfig configuration_switch_transformers.py 模型超参数,含专家数、容量、路由系数等
SwitchTransformersTop1Router modeling_switch_transformers.py#L52-L107 Top-1 路由:按容量约束为 token 选专家
SwitchTransformersExperts modeling_switch_transformers.py#L157-L176 8 个稠密专家(DenseActDense)的调度与聚合
SwitchTransformersSparseMLP modeling_switch_transformers.py#L179-L191 路由器 + 专家组的 MoE 前馈模块
router_z_loss_func / load_balancing_loss_func modeling_switch_transformers.py#L874-L930 训练时的 z-loss 与负载均衡辅助损失
SwitchTransformersForConditionalGeneration modeling_switch_transformers.py#L938-L1093 带 LM 头的 Seq2Seq 生成模型

2. 快速上手:加载检查点做序列到序列生成

官方文档给出的核心用法是预测被掩码的 token。以下代码完整继承自文档,可直接复制运行(前提:已安装 PyTorch,并 pip install transformers torch):

from transformers import AutoModelForSeq2SeqLM, AutoTokenizer


tokenizer = AutoTokenizer.from_pretrained("google/switch-base-8")
model = AutoModelForSeq2SeqLM.from_pretrained("google/switch-base-8", device_map="auto")

input_text = "The capital of France is <extra_id_0>."
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(0)

outputs = model.generate(input_ids)
print(tokenizer.decode(outputs[0]))

几点实操说明:

  • 输入文本中的 <extra_id_0> 是 Switch Transformers tokenizer 词表中的特殊"占位"token(其词表在 32k SentencePiece 基础上新增了 100 个 <extra_id_*> token,vocab_size 默认为 32128),用于提示模型"此处需要生成内容"。
  • device_map="auto" 依赖 accelerate 自动做设备映射;input_ids.to(0) 将输入放到第一个可用设备上,两者需保持一致。
  • AutoModel 外,也可以走 Pipeline 接口(文档同时提供了 Pipeline 与命令行两种入口),例如 pipeline("text2text-generation", model="google/switch-base-8", tokenizer="google/switch-base-8")

3. SwitchTransformersConfig:稀疏化参数全解

文档的 [[autodoc]] SwitchTransformersConfig 一节对应的完整参数定义在 configuration_switch_transformers.py#L54-L108。以下为源码中的全部默认值:

参数 默认值 含义
vocab_size 32128 词表大小(32000 SentencePiece + 100 个 extra_id + 特殊 token)
d_model 768 模型隐藏维度(hidden_size 的映射名)
d_kv 64 注意力中 key/value 的投影维度
d_ff 2048 前馈层中间维度
num_layers 12 编码器层数(num_hidden_layers 的映射名)
num_decoder_layers 12(None 时回退为 num_layers 解码器层数
num_heads 12 注意力头数
num_experts 8 每个稀疏层的专家数量
expert_capacity 64 单个专家每批 token 可服务的最大 token 数
num_sparse_encoder_layers 3 编码器中稀疏(MoE)FFN 层的数量
num_sparse_decoder_layers 3 解码器中稀疏(MoE)FFN 层的数量
router_bias False 路由分类器是否加 bias
router_jitter_noise 0.01 训练时给 token 输入注入的均匀噪声幅度
router_dtype "float32" 路由器计算精度,文档建议保持 float32(见第 4 节)
router_ignore_padding_tokens False 路由时是否忽略 padding token
relative_attention_num_buckets 32 相对位置偏置分桶数
relative_attention_max_distance 128 相对位置分桶的最大距离
dropout_rate 0.1 dropout 概率
layer_norm_epsilon 1e-6 层归一化数值稳定项
router_z_loss_coef 0.001 z-loss 系数
router_aux_loss_coef 0.001 负载均衡辅助损失系数
dense_act_fn "relu" FFN 激活函数,可选 "relu""gated-gelu"(后者用于 Switch Transformers v1.1)
initializer_factor 1.0 权重初始化缩放因子(主要用于测试)
is_encoder_decoder True 编码器-解码器结构标志
add_router_probs False 是否输出路由概率以计算辅助损失
use_cache / pad_token_id / eos_token_id / tie_word_embeddings 常规生成与对齐相关

配置类还通过 attribute_map 把 T5 风格的命名映射到标准命名:hidden_size → d_modelnum_attention_heads → num_headsnum_hidden_layers → num_layers(见 configuration_switch_transformers.py#L56)。

稀疏层的分布由 sparse step 决定。 __post_init__ 中通过整除计算得出步长:

if self.num_sparse_encoder_layers > 0:
    self.encoder_sparse_step = self.num_layers // self.num_sparse_encoder_layers
else:
    self.encoder_sparse_step = self.num_layers  # HACK: this will create 0 sparse layers

即默认 12 层编码器、3 个稀疏层时 encoder_sparse_step = 4,每 4 层插入一个 MoE 层。文档 docstring 也特别提醒了一个边界情况:当 num_sparse_encoder_layers=0num_layers=1 时,由于步长计算方式仍可能产生一个稀疏层,该边界情况在现有检查点中不会出现。解码器侧(decoder_sparse_step)逻辑完全对称。

4. Top-1 路由器:容量约束、抖动噪声与选择性精度

SwitchTransformersTop1Router 是整个 MoE 的核心,实现在 modeling_switch_transformers.py#L52-L107。其 docstring 明确说明:每个 token 只能选择一个专家,按路由概率排序后依次填充,直到专家的 expert_capacity 用尽;没有保证每个 token 都会被某个专家处理,也没有保证每个专家至少收到一个 token——这正是 Switch 论文中"溢出即丢弃"的容量机制。

前向过程分四步:

  1. 选择性精度(selective precision):路由计算强制转换到 router_dtype(默认 float32)。源码注释直接引用论文说明——保持 float32 是为了训练稳定性,因此配置项 router_dtype 不建议改成低精度。
  2. 训练期抖动(jitter):训练且 router_jitter_noise > 0 时,对 token 输入乘上 U(1-ε, 1+ε) 均匀噪声(默认 ε=0.01),用于缓解"专家偏好固化":
if self.training and self.jitter_noise > 0:
    hidden_states *= torch.empty_like(hidden_states).uniform_(1.0 - self.jitter_noise, 1.0 + self.jitter_noise)
  1. Top-1 选择与容量掩码:对 classifier(hidden_states) 做 softmax 后取 torch.max 得到每个 token 的首选专家,转成 one-hot,再用 torch.cumsum 沿序列维累计每个专家的"已到达 token 数",构造容量掩码:
expert_index = torch.nn.functional.one_hot(expert_index, num_classes=self.num_experts)
token_priority = torch.cumsum(expert_index, dim=-2)
# mask if the token routed to the expert will overflow
expert_capacity_mask = token_priority <= self.expert_capacity
expert_index = expert_index * expert_capacity_mask
  1. 输出:返回 (router_probs, expert_index, router_logits) 三元组,其中 router_probs 是首选专家的概率(shape 为 (batch, seq, 1),后续作为混合权重),router_logits 保留原始 logits 用于 z-loss。

5. 专家调度:SwitchTransformersSparseMLP 与 DenseActDense

SwitchTransformersExperts 继承自 nn.ModuleDict,持有 expert_0expert_{num_experts-1}num_experts 个专家,每个专家都是一个 SwitchTransformersDenseActDenseWi → act → Dropout → Wo,激活由 dense_act_fn 决定,默认 relu,无 bias 线性层,见 modeling_switch_transformers.py#L135-L154)。

前向调度逻辑(modeling_switch_transformers.py#L164-L176)是"按专家循环分发":

final_hidden_states = torch.zeros_like(hidden_states)
expert_mask = selected_experts.permute(2, 1, 0)

expert_hit = torch.greater(expert_mask.sum(dim=(-1, -2)), 0).nonzero()
for expert_idx in expert_hit:
    idx, top_x = torch.where(expert_mask[expert_idx].squeeze(0))
    current_state = hidden_states[None, top_x].reshape(-1, hidden_states.shape[-1])
    current_hidden_states = self[f"expert_{expert_idx[0]}"](current_state) * routing_weights[top_x, idx, None]
    final_hidden_states.index_add_(0, top_x, current_hidden_states.to(hidden_states.dtype))

几个要点:

  • 只有实际"被选中且未溢出"的专家会执行前向(expert_hit 过滤),这实现了稀疏计算:单 token 只过一个专家的 d_model → d_ff → d_model 两层投影。
  • 专家输出乘以路由权重 routing_weights 后用 index_add_ 累加回原位置——被容量掩码丢弃的 token 在 MoE 层得到零输出,依赖残差连接保留其原始表示。
  • 源码中 SwitchTransformersSparseMLP 注释标注 # inherit from mixtral,说明该文件由 modular_switch_transformers.py 自动生成(文件头部有明确的自动生成警告,修改须落在 modular 文件上)。

SwitchTransformersLayerFF 是 FFN 层的统一封装(modeling_switch_transformers.py#L194-L223):is_sparse=True 时用 SwitchTransformersSparseMLP,否则用稠密 SwitchTransformersDenseActDense,之后做 LayerNorm → MLP → Dropout → 残差。每个编码器/解码器块 SwitchTransformersBlock 按"自注意力(解码器块额外含交叉注意力)→ FFN"组装,稀疏与否由第 3 节的 sparse_step 规则在 SwitchTransformersStack.__init__ 中逐层判定(modeling_switch_transformers.py#L680-L690):

sparse_step = config.decoder_sparse_step if self.is_decoder else config.encoder_sparse_step
...
for i in range(config.num_layers):
    is_sparse = (i % sparse_step == 1 or sparse_step == 1) if sparse_step > 0 else False

注意该判定式意味着稀疏层从索引 1 开始、每隔 sparse_step 层出现一次,且 sparse_step == 1 时全部层都是稀疏层(测试配置即用此方式构造"全稀疏"模型)。

6. 训练侧:负载均衡损失与 z-loss

文档的 SwitchTransformersForConditionalGeneration 一节背后的关键工程点在于:只要前向传入 output_router_logits=True 且提供 labels,模型就会自动把 MoE 辅助损失加进总损失(modeling_switch_transformers.py#L1033-L1062):

if output_router_logits:
    ...
    encoder_z_loss = router_z_loss_func(encoder_router_logits)
    encoder_aux_loss = load_balancing_loss_func(encoder_router_probs, encoder_expert_indexes)
    ...
if labels is not None:
    loss = loss_fct(lm_logits.view(-1, ...), labels.view(-1))
    if output_router_logits:
        z_loss = self.router_z_loss_coef * (encoder_z_loss + decoder_z_loss)
        aux_loss = self.router_aux_loss_coef * (encoder_aux_loss + decoder_aux_loss)
        loss = loss + z_loss + aux_loss

两个损失函数的定义与公式:

  • z-lossmodeling_switch_transformers.py#L874-L891):z_loss = mean( logsumexp(logits)^2 )。docstring 说明它来自 Google 的《Designing Effective Sparse Expert Models》,用于约束路由 logits 幅度、提升训练稳定性。
  • 负载均衡辅助损失modeling_switch_transformers.py#L894-L930):实现论文公式 (4)-(6)。对每个专家,取"分到该专家的 token 占比"与"该专家的平均路由概率"的逐元素乘积,对专家取均值后乘以 num_experts^2 缩放,惩罚路由分布严重失衡:
tokens_per_group_and_expert = torch.mean(expert_mask, axis=-2)
router_prob_per_group_and_expert = torch.mean(router_probs, axis=-2)
return torch.mean(tokens_per_group_and_expert * router_prob_per_group_and_expert) * (num_experts**2)

若某侧 sparse_step <= 1 的判定不成立(源码以 encoder_sparse_step > 1 作为"存在稀疏层"的条件),对应侧损失直接置 0。输出对象 Seq2SeqMoEOutput 会同时返回 encoder_z_loss / encoder_aux_loss / decoder_z_loss / decoder_aux_loss / encoder_router_logits / decoder_router_logits,便于训练监控。

7. 架构细节:RMS 风格归一化、分桶相对位置偏置与共享词嵌入

在 T5 骨架之上,源码还有三处值得注意的实现细节:

  1. 只缩放不平移的 LayerNormSwitchTransformersLayerNormmodeling_switch_transformers.py#L110-L132)不 subtract mean、没有 bias,等价于 RMS Layer Normalization;方差在 fp32 下累加以保证半精度输入的数值稳定。
  2. 分桶相对位置偏置(T5 风格)SwitchTransformersAttention 用 32 个桶、最大距离 128 把相对位置映射为可学习偏置(relative_attention_num_buckets / relative_attention_max_distance,见 modeling_switch_transformers.py#L298-L361)。第一层的偏置在整条层栈中共享复用(Stack.forwardposition_bias = self_attention_position_bias 逐层传递)。另有一处与 T5 的差异:scaling = 1.0,即 QK 点积不做 1/sqrt(d) 缩放,相对位置偏置直接加在注意力分数上(源码注释 L278-L279)。
  3. 三处词嵌入共享SwitchTransformersForConditionalGenerationencoder.embed_tokensdecoder.embed_tokenslm_head 全部绑定到同一组 shared 权重(_tied_weights_keysmodeling_switch_transformers.py#L939-L943);且在 tie_word_embeddings=True 时,LM 头前会把隐藏态乘以 d_model^{-0.5} 做 rescale(modeling_switch_transformers.py#L1020-L1023)。

模型类层级为:SwitchTransformersPreTrainedModel(负责权重初始化,含路由器与专家各投影层的正态初始化)→ SwitchTransformersModel(无头 encoder+decoder,返回 Seq2SeqMoEModelOutput)→ SwitchTransformersForConditionalGeneration(加 LM 头与 MoE 损失)→ SwitchTransformersEncoderModel(仅编码器栈)。解码器支持 EncoderDecoderCache 增量 KV 缓存,generate 与第 2 节示例因此可直接工作。

8. 8-bit 量化实战:降低大模型的显存负担

官方文档第二块示例用 bitsandbytes 将权重压到 8-bit 以降低显存占用(完整继承自原文档):

# pip install bitsandbytes
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, BitsAndBytesConfig


tokenizer = AutoTokenizer.from_pretrained("google/switch-base-8")
quantization_config = BitsAndBytesConfig(load_in_8bit=True)
model = AutoModelForSeq2SeqLM.from_pretrained("google/switch-base-8", device_map="auto", quantization_config=quantization_config)

input_text = "The capital of France is <extra_id_0>."
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(0)

outputs = model.generate(input_ids)
print(tokenizer.decode(outputs[0]))

说明与适用前提:

  • BitsAndBytesConfig(load_in_8bit=True) 只量化权重为 8-bit(bitsandbytes 的 Linear8bitLt),激活与路由逻辑不受影响;若还需进一步压缩可参考 4-bit 配置(load_in_4bit + bnb_4bit_quant_type 等),更多后端见 量化总览bitsandbytes 量化指南
  • load_in_8bit 需要 CUDA 环境与已安装的 bitsandbytes;量化权重为 int8,源码中 SwitchTransformersDenseActDense.forwardwo.weight.dtype == torch.int8 做了跳过 dtype 对齐的分支(modeling_switch_transformers.py#L147-L153),保证量化权重可正常前向。
  • MoE 模型参数量大头在专家权重(num_experts × d_model × d_ff × 2),因此量化对显存收益显著。

9. 测试与验证

仓库自带的测试位于 tests/models/switch_transformers/test_modeling_switch_transformers.py,其中 SwitchTransformersModelTester 构造了一个微型配置(vocab_size=99d_model=32、2 层、expert_capacity=100router_jitter_noise=0.0)来验证稀疏层分布、生成、pipeline 等通用性质。直接验证路由与损失函数的写法可以参考测试文件的导入方式:

from transformers.models.switch_transformers.modeling_switch_transformers import (
    load_balancing_loss_func,
    router_z_loss_func,
)

此外,convert_big_switch.py 提供了从 Google 原始检查点到 Transformers 格式的转换脚本,可用于理解权重命名映射(例如原始 expert_0/...experts.expert_0.wi/wo 的对应关系)。

10. 小结

Switch Transformers 在 Transformers 中的实现是一个"完整可训练"的稀疏 MoE 参考:SwitchTransformersConfigexpert_capacitynum_expertssparse_step 三组参数控制稀疏度;SwitchTransformersTop1Router 用 cumsum 容量掩码实现论文中的 top-1 + 容量丢弃路由,并以 float32 选择性精度与 jitter 噪声保证训练稳定;SwitchTransformersExperts 以按专家分发 + index_add_ 聚合完成稀疏计算;router_z_loss_funcload_balancing_loss_func 则按系数自动并入总损失。推理侧只需 AutoModelForSeq2SeqLM.from_pretrained("google/switch-base-8") 即可生成,配合 BitsAndBytesConfig(load_in_8bit=True) 可进一步降低显存门槛。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
898
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
921
1.84 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.8 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
531
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
519
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
391