JetMoe 模型文档导读:在 Transformers 中使用混合注意力头与稀疏专家激活的 8B MoE 架构
JetMoe-8B 是一款由 Yikang Shen 与 MyShell 团队开发的 8B 规模混合专家(Mixture-of-Experts, MoE)语言模型,其设计目标是在有限预算下达到接近 LLaMA2 的性能。它在每个 Transformer 块中同时引入"注意力头混合(Mixture of Attention Heads, MoA)"与"MLP 专家混合(Mixture of MLP Experts, MoE)"两层稀疏激活结构,通过对每个 token 只激活部分专家来实现远高于同规模稠密模型的训练吞吐。本文基于 JetMoe 官方模型文档,并结合仓库内 configuration_jetmoe.py 与 modeling_jetmoe.py 的源码实现,完整讲解 JetMoeConfig 的每一项超参数、三类顶层 API(JetMoeModel / JetMoeForCausalLM / JetMoeForSequenceClassification)的用法,以及其稀疏路由、辅助负载均衡损失与注意力稀疏化的底层原理。读完本文,你将掌握如何加载 JetMoe 预训练权重、如何用自定义配置从零搭建模型、如何做因果语言建模与序列分类微调,以及如何解读模型输出的 router logits 与 aux_loss。
JetMoe 架构概述
根据官方文档,JetMoe 项目由 Yikang Shen 与 MyShell 开发,该论文于 2023-06-07 发布在 Hugging Face papers 上,并于 2024-05-14 合入 Transformers。
其核心设计是受到 ModuleFormer 启发的稀疏激活架构:
- 每个 JetMoe 块由两个 MoE 层组成:
Mixture of Attention Heads(注意力头混合,MoA)与Mixture of MLP Experts(MLP 专家混合,MoE)。这一点在源码中得到直接印证——modeling_jetmoe.py 中每个JetMoeDecoderLayer同时持有self.self_attention = JetMoeAttention(config, layer_idx)与self.mlp = JetMoeMoE(config),其中JetMoeAttention内部使用JetMoeMoA来生成多头查询。 - 稀疏激活:给定输入 token 后,模型只激活其中一部分专家来处理,从而大幅降低同批计算量。
- 训练吞吐收益:官方文档指出,JetMoe-8B 在 96 块 H100 GPU 集群上、配合直截了当的 3 路流水线并行(3-way pipeline parallelism)策略,训练吞吐约为每天 100B tokens。
因此 JetMoe 属于典型的"低成本训出可用模型"路线的代表实现,其注意力部分同样做了稀疏化,这与仅对 FFN 做专家混合的传统 MoE(如 Mixtral)有明显区别。
JetMoeConfig 全参数详解
JetMoeConfig 在 configuration_jetmoe.py 中定义,model_type = "jetmoe"。它继承 PreTrainedConfig,通过 @strict 与 @auto_docstring(checkpoint="jetmoe/jetmoe-8b") 装饰,前者启用严格的类型与架构校验,后者在生成 API 文档时自动代入 jetmoe/jetmoe-8b 这一默认 checkpoint。下表汇总该配置类的全部字段及其默认值,均以当前仓库源码为准:
| 配置字段 | 默认值 | 说明 |
|---|---|---|
vocab_size |
32000 | 词表大小 |
hidden_size |
2048 | 隐藏层维度 |
num_hidden_layers |
12 | Transformer 解码层数量 |
num_key_value_heads |
16 | KV 头数量 |
kv_channels |
128 | key/value 张量的通道数,即每个注意力头的维度(head_dim) |
intermediate_size |
5632 | MLP 专家内部的隐藏维度 |
max_position_embeddings |
4096 | 最大序列长度 |
activation_function |
"silu" | MLP 激活函数,JetMoeMoE 通过 ACT2FN[config.activation_function] 查表加载,并实现 SiLU 门控乘法 |
num_local_experts |
8 | MoE 与 MoA 中的专家数量 |
num_experts_per_tok |
2 | 每个 token 路由到的 top-k 专家数量 |
output_router_logits |
False | 是否输出 router logits(用于计算辅助负载均衡损失) |
aux_loss_coef |
0.01 | 辅助负载均衡损失的系数 |
use_cache |
True | 推理时是否使用 KV cache |
bos_token_id / eos_token_id |
1 / 2 | 起止符 token id |
pad_token_id |
None | 填充符 token id |
tie_word_embeddings |
True | 是否绑定输入输出词嵌入 |
rope_parameters |
None | RoPE 旋转位置编码参数(如 rope_type、rope_theta),传给 JetMoeRotaryEmbedding |
rms_norm_eps |
1e-6 | RMSNorm 的 epsilon |
initializer_range |
0.01 | 专家参数初始化标准差 |
attention_dropout |
0.0 | 注意力 dropout 概率 |
在 __post_init__ 中有两个值得注意的派生逻辑:
- 注意力头数量由配置推导:
self.num_attention_heads = self.num_key_value_heads * self.num_experts_per_tok,即默认配置下为16 × 2 = 32个注意力头。这意味着每个 token 实际会使用 top-k 个"注意力专家",每个专家贡献num_key_value_heads个头,因此总头数是两者的乘积。 - 架构严格校验:
validate_architecture()在num_experts_per_tok > num_local_experts时抛出ValueError,保证路由的 top-k 不会超过专家总数。
此外该类设置了 keys_to_ignore_at_inference = ["past_key_values"] 与 attribute_map = {"head_dim": "kv_channels"}。后者说明在 JetMoe 中 head_dim 只是 kv_channels 的别名。
快速上手:初始化配置与模型
沿用源码 docstring 中的最小示例,你可以直接基于默认配置初始化一个模型:
from transformers import JetMoeModel, JetMoeConfig
# 初始化一个 JetMoe 风格(4B 量级)的配置
configuration = JetMoeConfig()
# 从配置初始化模型
model = JetMoeModel(configuration)
# 读取模型的实际配置
configuration = model.config
当使用官方预训练权重时,请直接使用 AutoModelForCausalLM 系列按模型 id jetmoe/jetmoe-8b 加载(当前仓库标注的默认 checkpoint 即为该权重),例如:
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("jetmoe/jetmoe-8b")
tokenizer = AutoTokenizer.from_pretrained("jetmoe/jetmoe-8b")
仓库通过 models/jetmoe 目录 下的 __init__.py 使用 _LazyModule 实现懒加载,JetMoeConfig、JetMoeModel、JetMoeForCausalLM 等均可从 transformers 顶层直接导入,无需关心内部模块路径。
模块类与底层实现解读
本节将官方文档列出的三个 autodoc 类与源码实现一一对应,并揭示各关键机制在代码中的位置。
JetMoeModel:基础解码器栈
JetMoeModel(modeling_jetmoe.py)是无输出头的基础模型,结构与 LLaMA 类模型接近:embed_tokens 词嵌入 → nn.ModuleList 堆叠的 JetMoeDecoderLayer → 末端 JetMoeRMSNorm,外加全局共享的 JetMoeRotaryEmbedding。其 forward 输出类型为 MoeModelOutputWithPast,包含 last_hidden_state 与 past_key_values,注释明确说明它与 Mistral 的唯一差异就是输出类型为 MoE 专用结构。模型本身并不对 router_logits 做任何损失计算,只负责在 output_router_logits=True 时透传各层路由结果。
JetMoeForCausalLM:因果语言建模与辅助损失
JetMoeForCausalLM(modeling_jetmoe.py)在 JetMoeModel 之上叠加 lm_head 线性层,并实现了 GenerationMixin,因此可以直接用于文本生成、继续预训练与指令微调。值得注意的实现细节:
- 词嵌入绑定:
_tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"},与配置中tie_word_embeddings=True对应。 - 辅助负载均衡损失:
load_balancing_loss_func(modeling_jetmoe.py)实现了 Switch Transformer 论文(论文链接)式方程 (4)–(6) 的损失函数,用于惩罚专家路由过于不均衡。它按层累积"每专家分配的 token 数"与"每专家路由概率之和",最后计算sum(tokens_per_expert * router_prob_per_expert) * num_experts;传入了attention_mask时,会用扁平化 mask 对 padding token 加权排除,内存占用保持在O(seq_len * num_experts)量级,与层数无关。 - 损失组成:仅在
output_router_logits=True时才计算 aux_loss;若同时提供了labels,总损失为loss += self.aux_loss_coef * aux_loss,其中aux_loss_coef默认 0.01。 - 生成友好:通过
logits_to_keep参数可以只计算最后 N 个位置的 logits,避免为长序列浪费算力。
一个包含 MoE 辅助损失的最小训练式前向示例:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("jetmoe/jetmoe-8b")
tokenizer = AutoTokenizer.from_pretrained("jetmoe/jetmoe-8b")
inputs = tokenizer("Machine learning is", return_tensors="pt")
labels = inputs["input_ids"]
outputs = model(**inputs, labels=labels, output_router_logits=True)
loss = outputs.loss # CE loss(含 aux_loss_coef * aux_loss)
aux_loss = outputs.aux_loss # 可单独观察的路由均衡损失
router_logits = outputs.router_logits # 各层 router 的 logits 元组
注意:由于 MoE 路由训练易发生"专家坍缩",官方建议在预训练/微调阶段开启 output_router_logits 以将辅助损失计入反向传播。默认 aux_loss_coef=0.01 属于较轻的约束强度,如需更强均衡可适当调大。
JetMoeForSequenceClassification:序列分类头
JetMoeForSequenceClassification(modeling_jetmoe.py)直接复用通用的 GenericForSequenceClassification 与 JetMoePreTrainedModel 组合生成,因此拥有与其它模型一致的 num_labels 分类头、problem_type 自动推断与 loss 返回行为,适用于文本分类、情感分析等下游任务:
from transformers import AutoModelForSequenceClassification
model = AutoModelForSequenceClassification.from_pretrained("jetmoe/jetmoe-8b", num_labels=2)
稀疏路由与注意力混合的实现机制
要真正理解 JetMoe,需要把目光落到源码中的几个核心模块上。
JetMoeTopKGating:路由决策
JetMoeTopKGating 是一个无 bias 的 nn.Linear(input_size, num_experts)。它在 forward 中依次完成:对每个 token 计算 logits → 取 top_k 的 logits 与索引 → 对 top-k logits 做 softmax 得到门控权重 → 统计每个专家分到的 token 数 expert_size → 按专家分组排序 token 并回传对应的批量索引与门控值。源码注释专门提醒:expert_size.tolist() 这一数据依赖操作会导致 torch.compile 的 fullgraph 模式失败,因此基类中明确设置了 _can_compile_fullgraph = False。
JetMoeMoE:稀疏 FFN 专家
JetMoeMoE 是标准的稀疏门控专家层:JetMoeParallelExperts 以 [num_experts, ...] 的权重布局并行存储全部专家参数,输入按路由分组后经 input_linear 升维并切分为两份,执行 activation(x0) * x1(即 SiLU 门控乘法的等价实现),再经 output_linear 还原;最后用门控值加权,通过 index_add 把各专家的输出加回到原始 token 顺序,并追加可学习的 bias。
JetMoeMoA:注意力头混合(稀疏注意力)
JetMoeMoA 是 JetMoe 区别于传统 MoE 的核心创新:把"对专家做投影再按路由加权合并"的思想从 FFN 搬到了注意力层。它的每个专家由一对 query-与 output-projection 组成,采用 map/reduce 两阶段设计:
map:用 router 决定 top-k 注意力专家,把 token 分组后只计算被选中专家的 query 投影,再恢复为[batch, seq_len, top_k, hidden]的形状;reduce:按同样的拓扑信息(topo_info)把输出投影做在专家内部,乘上门控权重后把 top-k 路结果index_add合并回原顺序。
JetMoeAttention(modeling_jetmoe.py)在此基础上共享地投影 K/V(kv_proj),随后按 top_k 复制 K/V,与各注意力专家生成的 Q 计算注意力。代码注释明确写了一句区别于其它模型的重要提示:"This is different from other models where we repeat k/v heads instead of repeat interleaving them"——即这里 K/V 是按 torch.repeat(连续复制块)而非逐头插值(repeat_interleave)展开的,理解这一点对自行调试注意力形状很有帮助。
注意力实现选择与稀疏输出
JetMoe 属于"稀疏激活、稠密注意力计算"的 MoE 家族成员:虽然每个 token 只经过 top-k 专家,但其注意力 logits 仍需在所有展开后的头之间计算。仓库实现通过 ALL_ATTENTION_FUNCTIONS 接口动态选择后端:
attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
self.config._attn_implementation, eager_attention_forward
)
也就是说你可以用标准的 attn_implementation="eager" | "sdpa" | "flash_attention_2" 等开关选择注意力后端。官方模型文档顶部为此标注了 FlashAttention 与 SDPA 的徽标,源码中 JetMoePreTrainedModel 也声明了 _supports_flash_attn = True、_supports_sdpa = True、_supports_flex_attn = True 三个能力位。默认的 eager_attention_forward(modeling_jetmoe.py)在 fp32 下做 softmax 以保证数值稳定。
从源码结构中得出的注意事项
- 生成与缓存:
JetMoeAttention在未传入layer_idx时会告警,若使用缓存则必须保证每层携带正确的layer_idx;基类把past_key_values排除在设备放置之外(_skip_keys_device_placement),多卡场景无需手动搬迁缓存。 - 前向可记录输出:
_can_record_outputs表明你可以通过统一的输出捕获机制分别记录router_logits(来自JetMoeAttention与JetMoeTopKGating)、各JetMoeDecoderLayer的hidden_states与attentions,方便做路由可视化与调试。 - 测试验证:仓库的 test_modeling_jetmoe.py 覆盖了配置校验、前向输出形状、路由 logits、辅助损失以及与其他 MoE 模型一致的通用行为,是复现本文所述行为的权威参照。
- 实现来源:modeling_jetmoe.py 顶部声明该文件由 modular_jetmoe.py 自动生成,若需给 JetMoe 贡献自定义改造应改在 modular 源文件上,并受 CI 的"生成文件与 modular 一致性"检查约束。
综上,JetMoe 通过"注意力头混合 + MLP 专家混合"的双重稀疏化,在压缩激活计算的同时保留了接近稠密模型的质量;而 output_router_logits、aux_loss_coef 与配套的 MoeCausalLMOutputWithPast 输出结构,则为在 Transformers 生态内做负载均衡训练与下游微调提供了开箱即用的支持。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0627
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00