首页
/ Transformers 中的 Mellum:JetBrains 代码向 MoE 大模型架构解析与使用指南

Transformers 中的 Mellum:JetBrains 代码向 MoE 大模型架构解析与使用指南

2026-09-07 20:05:48作者:郁楠烈Hubert

导读

Mellum 是 JetBrains 推出的代码领域专用 Mixture-of-Experts(MoE)语言模型,其实现以 Qwen3-MoE 架构为基底,引入逐层类型(per-layer-type)RoPE 与交错滑动窗口注意力。本文以 Mellum 模型文档 为主体,结合其配置与建模源码,讲解如何在 transformers 中加载 Mellum 做代码补全/生成,并深入剖析 64 专家 Top-8 路由、28 层稀疏架构、QK-Norm 与双 RoPE 频段的底层实现,读完你既能直接跑通推理示例,也能读懂配置中的每个关键参数。

一、Mellum 是什么:一次来自 Qwen3-MoE 与 Laguna 的架构组合

Mellum 由 JetBrains 贡献给 Hugging Face Transformers(日期为 2026-05-28,见 mellum.md 首部说明),定位是代码聚焦(code-focused)的 MoE 语言模型。原文档给出的核心事实是:

  • 总参数量 12B,每个 token 仅激活 2.5B 参数;
  • 28 层,每层含 64 个路由专家(routed experts),每 token 激活其中 8 个
  • 架构承袭 Qwen3-MoE,并叠加了逐层类型的 RoPE(per-layer-type RoPE)与交错滑动窗口注意力(interleaved sliding window attention)。

从源码看,这份“继承”在仓库中是显式的。Mellum 代码采用 modular 构建方式,其源头文件 modular_mellum.py 直接声明:

  • MellumConfig 继承 Qwen3MoeConfig
  • 注意力、MoE 块等复用 qwen3_moe 模块中的 Qwen3MoeAttentionQwen3MoeSparseMoeBlock 等类;
  • 同时组合了 laguna 模块中的 LagunaDecoderLayerLagunaRotaryEmbedding

也就是说,Mellum 可以理解为 "Qwen3-MoE 骨干 + 混合逐层结构 + 双层 RoPE 配置" 的代码模型变体。模型文档页对应的代码支持集中在 src/transformers/models/mellum/,含 4 个文件:__init__.pyconfiguration_mellum.pymodeling_mellum.pymodular_mellum.py(其中 configuration_*modeling_* 均由 modular 文件自动生成,编辑需作用于 modular 源文件)。

从代码结构推断,本仓库版本对应的官方权重标识为 JetBrains/Mellum2-12B-A2.5B-Base(出现在 @auto_docstring(checkpoint=...) 装饰器与模型文档示例中)。是否需要额外的词元器(tokenizer)配合,以实际 Hub 权重为准。

二、快速开始:两种方式驱动 Mellum 生成代码

原文档给出了两条可复制、可运行的生成路径,二者等价,按使用习惯任选其一即可。使用前请先安装 transformers 及其依赖(pip install transformers / pip install -e .),并保证有足够显存加载模型

2.1 方式一:通过 pipeline(最少代码)

Mellum 支持文本生成任务,可直接用高层 Pipeline 接口:

from transformers import pipeline

pipe = pipeline(
    task="text-generation",
    model="JetBrains/Mellum2-12B-A2.5B-Base",
)
pipe("def fibonacci(n):")

pipeline 会自动完成 tokenizer 与模型的加载,适合快速验证模型效果;代码模型通常会补全出后续的 Python 函数体。

2.2 方式二:通过 AutoModelForCausalLM(更精细的控制)

需要显式拿到 logits、控制采样参数或逐段生成时,推荐使用 Auto 类手动组合 tokenizer 与模型:

from transformers import AutoModelForCausalLM, AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("JetBrains/Mellum2-12B-A2.5B-Base")
model = AutoModelForCausalLM.from_pretrained(
    "JetBrains/Mellum2-12B-A2.5B-Base",
    device_map="auto",          # 自动分配层到可用设备,多 GPU / 混合 CPU-GPU 均可用
)
input_ids = tokenizer("def fibonacci(n):", return_tensors="pt").to(model.device)

output = model.generate(**input_ids, max_new_tokens=50)
print(tokenizer.decode(output[0], skip_special_tokens=True))

要点:

  • device_map="auto" 依赖 accelerate,可自动把 12B 总参数分散到多卡;
  • model.generate 来自 MellumForCausalLM 所继承的 GenerationMixin(见 modeling_mellum.py),因此 transformers 提供的 max_new_tokenstemperaturetop_p 等解码参数均适用;
  • Mellum 的 MellumForCausalLM 在输入 labels 时也会返回 MoE 相关的 aux_lossloss,与常规因果语言模型用法一致(详细见下文第七节)。

三、MellumConfig:逐参数解读 12B MoE 配置

MellumConfigconfiguration_mellum.py)是模型的唯一配置入口,model_type = "mellum"。下表汇总了默认配置参数,与 12B / 2.5B 激活的规格一一对应。

3.1 基础尺寸参数

参数 默认值 含义
vocab_size 98304 词表大小
hidden_size 2304 隐藏层宽度
intermediate_size 7168 稠密层(dense MLP) 的中间维度
num_hidden_layers 28 解码器层数
num_attention_heads 32 注意力头数
num_key_value_heads 4 KV 头数(GQA 分组,每组 32/4=8 个 query 头)
head_dim 128 每头维度,注意与 hidden_size / heads = 72 不同,显式独立配置
hidden_act "silu" 激活函数(专家与稠密 MLP 共用,即 SiLU/SwiGLU 门控)
max_position_embeddings 131072 最大位置数(128K 语境)
initializer_range 0.02 权重初始化标准差
rms_norm_eps 1e-6 RMSNorm 的 epsilon
attention_bias False 注意力投影是否带偏置
attention_dropout 0.0 注意力 dropout 率
sliding_window 1024 滑动窗口注意力层的窗口大小
tie_word_embeddings False 输入/输出嵌入是否共享
use_cache True 默认启用 KV Cache

3.2 MoE 相关参数

参数 默认值 含义
num_experts 64 每层路由专家总数
num_experts_per_tok 8 每 token 激活的专家数(Top-8)
moe_intermediate_size 896 单个专家的 FFN 中间维度
norm_topk_prob True 对 Top-8 路由概率做归一化后再加权
output_router_logits False 前向时是否输出各层 router logits
router_aux_loss_coef 0.001 负载均衡辅助损失的权重系数

注意配置类中的 attribute_map = {"num_experts": "num_local_experts"},即旧属性名 num_local_experts 会被自动映射到 num_experts,便于兼容 MoE 通用序列化。

3.3 逐层结构参数(Mellum 的差异化核心)

  • layer_types(默认 None:每层注意力类型,取值 "full_attention""sliding_attention",长度必须等于 num_hidden_layers。若为 None__post_init__ 会填为 28 个 "full_attention"(全量注意力)。实际权重的“交错”排布由 Hub 上的 config 显式给出。
  • mlp_layer_types(默认 None:每层 MLP 类型,取值 "dense""sparse",长度同为层数;None 时全部默认 "sparse",即默认逐层走 MoE。这是稠密/稀疏混合结构的开关,解码器构造时按 config.mlp_layer_types[layer_idx] 二选一(见 modeling_mellum.py)。

3.4 逐层 RoPE 参数与滑动窗口的关系

Mellum 的另一个特征是把 RoPE 参数按层类型分别配置(配置在 rope_parameters 字段,同时兼容 dict 与新的 RopeParameters 类型)。默认值如下(见 configuration_mellum.py):

层类型 rope_type rope_theta
full_attention "default" 500000.0
sliding_attention "default" 10000.0

即:全量注意力层使用高频基 500000(长距离依赖保持能力强),滑动窗口层使用低频基 10000——窗口内只有局部交互,不需要过大的 theta 来维持远端分辨率。__post_init__ 在触发 RoPE 校验时显式忽略 {"sliding_attention", "full_attention"} 这两个键,避免对新模型做旧式 rope_scaling 兼容校验;同时 convert_rope_params_to_dict 直接原样返回参数,因为新模型不存在旧格式字段。

除字段外,配置类还预置了 keys_to_ignore_at_inference = ["past_key_values"],并在 base_model_tp_plan / base_model_ep_plan / base_model_pp_plan 中声明了张量并行(TP)、专家并行(EP)与流水线并行(PP)的默认切分计划,说明 Mellum 从设计上就为分布式推理/训练做好了映射(详见第六节)。

四、从源码理解 Mellum 的逐层执行

把 28 层“拼装”起来的核心是 MellumModelmodeling_mellum.py)与 MellumDecoderLayer(同文件 L374)。其前向流程与常规解码器一致:embed_tokens → 逐层 Self-Attention + MLP 残差 → norm,差别集中在三处。

4.1 按层类型缓存 mask 与位置编码

MellumModel.forward 不是给每层单独算 mask,而是先扫描 set(config.layer_types) 中出现过的层类型,每种类型只构造一份掩码与一份 cos/sin,再按 layer_types[i] 索引分发:

  • 掩码侧:create_causal_mask 用于 "full_attention"create_sliding_window_causal_mask 用于 "sliding_attention"modeling_mellum.py);
  • 位置编码侧:self.rotary_emb(...) 对每个层类型分别计算(同文件 L518-L520),再由每层取 position_embeddings[self.config.layer_types[i]] 传入(L526)。

4.2 MellumRotaryEmbedding:一组按层类型分离的缓冲

MellumRotaryEmbeddingmodeling_mellum.py)在初始化时遍历层类型,为每种类型独立保存:

  • {layer_type}_inv_freq:逆频(nn.Bufferpersistent=False);
  • {layer_type}_original_inv_freq:原始逆频副本(供动态 RoPE 更新);
  • {layer_type}_attention_scaling:注意力缩放因子。

forward(x, position_ids, layer_type) 依据 layer_type 取用对应缓冲,在 fp32 下计算 cos/sin 后回投到输入 dtype。值得一提的实现细节(注释 # key difference to gemma3: partial rope):compute_default_rope_parameters 支持 partial_rotary_factor(默认 1.0),可将 RoPE 只施加到 head_dim 的前若干维,即部分旋转;逆频公式为标准的 1 / (base ** (arange(0, dim, 2) / dim)),基频 baserope_parameters[layer_type]["rope_theta"] 给出(同文件 L86-L96)。

4.3 注意力:QK-Norm + 滑动窗口 + 可插拔后端

MellumAttentionmodeling_mellum.py)在 Qwen3-MoE 注意力基础上叠加了逐头 QK-Norm

self.q_norm = MellumRMSNorm(self.head_dim, eps=config.rms_norm_eps)  # 仅作用在 head_dim 上
self.k_norm = MellumRMSNorm(self.head_dim, eps=config.rms_norm_eps)

q_proj/k_proj 输出被 reshape 为 (..., head_dim) 后先经过 q_norm/k_norm 再做 RoPE,可显著稳定训练。MellumRMSNorm 与 T5LayerNorm 等价(对输入做方差归一,无均值平移)。窗口逻辑由 sliding_window 字段承载:

self.sliding_window = config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None

当前层为滑动层时,该值被透传给注意力函数(注释 # diff with Llama 标出了与 Llama 的差异点);最终注意力实现通过 ALL_ATTENTION_FUNCTIONS.get_interface(config._attn_implementation, eager_attention_forward) 分发到 Eager / SDPA / FlashAttention / FlexAttention 后端。

4.4 MoE 块:3D 参数、Top-K 门控与辅助损失

  • MellumTopKRouter(同文件 L323):以 weight (64, 2304) 线性打分 → softmax(fp32)→ topk(8) → 若 norm_topk_prob 为真则对 Top-8 概率归一化,返回 (router_logits, router_scores, router_indices)
  • MellumExperts(同文件 L283):专家权重以 3D 张量存储——gate_up_proj (64, 2*896, 2304)down_proj (64, 2304, 896),便于按 experts 索引整体取出;每个命中专家上执行 silu(gate) * up 后投影回隐藏维,并按 top_k_weights 加权累加。MellumExperts@use_experts_implementation 装饰,可在集成层替换为融合 kernel(如 grouped GEMM)。
  • load_balancing_loss_func(同文件 L540):实现 Switch Transformer 式负载均衡损失(论文见 papers/2101.03961),统计每个专家的 token 分配占比与路由概率占比并求内积;实现上逐层累积后再归一化,使峰值内存保持 O(seq_len * num_experts) 而与层数无关,同时支持用 attention_mask 剔除 padding 的影响。

MellumForCausalLM.forwardoutput_router_logits=True(或 config 置真)时收集 router_logits 并计算 aux_loss;若同时给了 labels,则按 loss += router_aux_loss_coef * aux_loss 叠加到交叉熵上(modeling_mellum.py)。

4.5 稠密层与残差结构

非稀疏层由 MellumMLP 承担:gate_proj/up_proj/down_proj 三段式 SwiGLU,中间维度取 config.intermediate_size(7168)。每个 MellumDecoderLayer 结构为:input_layernorm → attention → 残差 → post_attention_layernormmlp(稀疏或稠密)→ 残差,并继承自 GradientCheckpointingLayer 以支持梯度检查点。

五、推理注意力后端与能力开关

MellumPreTrainedModel 类属性(modeling_mellum.py)可以确认模型对多种注意力实现的原生支持:

_supports_flash_attn = True
_supports_sdpa = True
_supports_flex_attn = True
_can_compile_fullgraph = True
_supports_attention_backend = True

这与模型文档首页的两枚能力徽章(FlashAttention、SDPA,见 mellum.md)吻合。你可以通过 model = AutoModelForCausalLM.from_pretrained(..., attn_implementation="flash_attention_2")"sdpa" / "eager" 显式选择后端(需要对应环境安装 flash-attn 等依赖);不指定时按环境自动择优。此外 _no_split_modules = ["MellumDecoderLayer"]_skip_keys_device_placement = ["past_key_values"] 配合 device_map/accelerate 切分,且默认使用 DynamicCache(MellumModel 前向中 DynamicCache(config=self.config))做 KV 缓存。

六、面向并行的默认切分计划(配置即架构设计)

MellumConfig 内置了三种分布式并行策略的默认“蓝图”,从侧面印证其面向大规模部署的设计:

  • TP 计划(base_model_tp_planq_proj/k_proj/v_projcolwiseo_projrowwiseq_norm/k_normreplicated_with_grad_allreduce;专家 gate_up_projpacked_colwisedown_projrowwise,且 layers.*.mlp.experts 整体标记为 moe_tp_experts;稠密 gate/up/down_proj 也按列/行切分。
  • EP 计划(base_model_ep_plan:仅切分 MoE 专家,注意力保持不切分(注释说明由 FSDP2 负责注意力权重分发),并注明“EP 规模可突破 num_kv_heads 限制”——这意味着专家并行不必被 4 个 KV 头卡住,扩卡更自由。
  • PP 计划(base_model_pp_planembed_tokens → layers → norm → lm_head 的段间输入输出契约,便于按层切流水线。

这些计划与类上的 _tp_plan_pp_plan_fsdp_plan(如 lm_head: "colwise_gather_output""keep_full_weight")共同构成 Mellum 在 TP/EP/PP/FSDP 多种并行模式下的默认行为,普通推理用户无需改动即可享用。

七、API 一览:MellumConfig / MellumModel / MellumForCausalLM

对应模型文档的 [[autodoc]] 区段,可用的核心符号为:

用途 关键点
MellumConfig 配置类 见第三节参数表;from_pretrained 会自动读取权重目录的 config.json
MellumModel 无头的基础解码器 forward(input_ids, attention_mask, position_ids, past_key_values, inputs_embeds, use_cache),返回 MoeModelOutputWithPastlast_hidden_state 等)
MellumForCausalLM 因果语言建模头 额外接受 labelsoutput_router_logitslogits_to_keep;返回 MoeCausalLMOutputWithPast,含 lossaux_lossrouter_logits

几个参数在代码中的行为:

  • position_ids:未显式传入时,代码会用 past_seen_tokens 与序列长度自动补齐(modeling_mellum.py);
  • logits_to_keep:仅对末尾 N 个位置计算 logits 以省显存,配合 labels 做自回归训练切片使用;
  • 初始化(_init_weights)对 MellumExpertsgate_up_proj/down_projMellumTopKRouter.weight 单独做了均值为 0、std=initializer_range 的正态初始化,并对旋转编码按层类型刷新逆频缓冲。

若要在没有官方权重时先本地跑通结构,可以直接用随机配置构建:

from transformers import MellumModel, MellumConfig

configuration = MellumConfig()          # 12B 全默认参数
model = MellumModel(configuration)      # 随机初始化,验证前向/结构

注意该示例会创建约 12B 参数的随机模型,仅适合小批量结构自检,勿在低资源环境下盲目执行。

八、延伸阅读路径

若要继续深入,建议在仓库内按以下顺序阅读:

综上,Mellum 以 Qwen3-MoE 为骨架、以“逐层类型”为编排手段,把稠密/稀疏 MLP、全量/滑动注意力与双频段 RoPE 组合进 28 层结构中,在 12B 总参数下维持 2.5B 激活的高效推理。对于代码生成、补全类场景,你可以基于上文任一示例直接替换模型标识体验;若要二次开发或微调,则需要从 modular_mellum.py 出发修改,并让改动同步到生成的建模与配置文件。

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

项目优选

收起
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