首页
/ JetMoe 模型文档导读:在 Transformers 中使用混合注意力头与稀疏专家激活的 8B MoE 架构

JetMoe 模型文档导读:在 Transformers 中使用混合注意力头与稀疏专家激活的 8B MoE 架构

2026-09-07 18:09:35作者:尤辰城Agatha

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.pymodeling_jetmoe.py 的源码实现,完整讲解 JetMoeConfig 的每一项超参数、三类顶层 API(JetMoeModel / JetMoeForCausalLM / JetMoeForSequenceClassification)的用法,以及其稀疏路由、辅助负载均衡损失与注意力稀疏化的底层原理。读完本文,你将掌握如何加载 JetMoe 预训练权重、如何用自定义配置从零搭建模型、如何做因果语言建模与序列分类微调,以及如何解读模型输出的 router logits 与 aux_loss。

JetMoe 架构概述

根据官方文档,JetMoe 项目由 Yikang ShenMyShell 开发,该论文于 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 全参数详解

JetMoeConfigconfiguration_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_typerope_theta),传给 JetMoeRotaryEmbedding
rms_norm_eps 1e-6 RMSNorm 的 epsilon
initializer_range 0.01 专家参数初始化标准差
attention_dropout 0.0 注意力 dropout 概率

__post_init__ 中有两个值得注意的派生逻辑:

  1. 注意力头数量由配置推导self.num_attention_heads = self.num_key_value_heads * self.num_experts_per_tok,即默认配置下为 16 × 2 = 32 个注意力头。这意味着每个 token 实际会使用 top-k 个"注意力专家",每个专家贡献 num_key_value_heads 个头,因此总头数是两者的乘积。
  2. 架构严格校验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 实现懒加载,JetMoeConfigJetMoeModelJetMoeForCausalLM 等均可从 transformers 顶层直接导入,无需关心内部模块路径。

模块类与底层实现解读

本节将官方文档列出的三个 autodoc 类与源码实现一一对应,并揭示各关键机制在代码中的位置。

JetMoeModel:基础解码器栈

JetMoeModelmodeling_jetmoe.py)是无输出头的基础模型,结构与 LLaMA 类模型接近:embed_tokens 词嵌入 → nn.ModuleList 堆叠的 JetMoeDecoderLayer → 末端 JetMoeRMSNorm,外加全局共享的 JetMoeRotaryEmbedding。其 forward 输出类型为 MoeModelOutputWithPast,包含 last_hidden_statepast_key_values,注释明确说明它与 Mistral 的唯一差异就是输出类型为 MoE 专用结构。模型本身并不对 router_logits 做任何损失计算,只负责在 output_router_logits=True 时透传各层路由结果。

JetMoeForCausalLM:因果语言建模与辅助损失

JetMoeForCausalLMmodeling_jetmoe.py)在 JetMoeModel 之上叠加 lm_head 线性层,并实现了 GenerationMixin,因此可以直接用于文本生成、继续预训练与指令微调。值得注意的实现细节:

  • 词嵌入绑定_tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"},与配置中 tie_word_embeddings=True 对应。
  • 辅助负载均衡损失load_balancing_loss_funcmodeling_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:序列分类头

JetMoeForSequenceClassificationmodeling_jetmoe.py)直接复用通用的 GenericForSequenceClassificationJetMoePreTrainedModel 组合生成,因此拥有与其它模型一致的 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 合并回原顺序。

JetMoeAttentionmodeling_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_forwardmodeling_jetmoe.py)在 fp32 下做 softmax 以保证数值稳定。

从源码结构中得出的注意事项

  • 生成与缓存JetMoeAttention 在未传入 layer_idx 时会告警,若使用缓存则必须保证每层携带正确的 layer_idx;基类把 past_key_values 排除在设备放置之外(_skip_keys_device_placement),多卡场景无需手动搬迁缓存。
  • 前向可记录输出_can_record_outputs 表明你可以通过统一的输出捕获机制分别记录 router_logits(来自 JetMoeAttentionJetMoeTopKGating)、各 JetMoeDecoderLayerhidden_statesattentions,方便做路由可视化与调试。
  • 测试验证:仓库的 test_modeling_jetmoe.py 覆盖了配置校验、前向输出形状、路由 logits、辅助损失以及与其他 MoE 模型一致的通用行为,是复现本文所述行为的权威参照。
  • 实现来源modeling_jetmoe.py 顶部声明该文件由 modular_jetmoe.py 自动生成,若需给 JetMoe 贡献自定义改造应改在 modular 源文件上,并受 CI 的"生成文件与 modular 一致性"检查约束。

综上,JetMoe 通过"注意力头混合 + MLP 专家混合"的双重稀疏化,在压缩激活计算的同时保留了接近稠密模型的质量;而 output_router_logitsaux_loss_coef 与配套的 MoeCausalLMOutputWithPast 输出结构,则为在 Transformers 生态内做负载均衡训练与下游微调提供了开箱即用的支持。

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

项目优选

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