Transformers 中的 Mellum:JetBrains 代码向 MoE 大模型架构解析与使用指南
导读
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模块中的Qwen3MoeAttention、Qwen3MoeSparseMoeBlock等类; - 同时组合了
laguna模块中的LagunaDecoderLayer、LagunaRotaryEmbedding。
也就是说,Mellum 可以理解为 "Qwen3-MoE 骨干 + 混合逐层结构 + 双层 RoPE 配置" 的代码模型变体。模型文档页对应的代码支持集中在 src/transformers/models/mellum/,含 4 个文件:__init__.py、configuration_mellum.py、modeling_mellum.py 与 modular_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_tokens、temperature、top_p等解码参数均适用;- Mellum 的
MellumForCausalLM在输入labels时也会返回 MoE 相关的aux_loss与loss,与常规因果语言模型用法一致(详细见下文第七节)。
三、MellumConfig:逐参数解读 12B MoE 配置
MellumConfig(configuration_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 层“拼装”起来的核心是 MellumModel(modeling_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:一组按层类型分离的缓冲
MellumRotaryEmbedding(modeling_mellum.py)在初始化时遍历层类型,为每种类型独立保存:
{layer_type}_inv_freq:逆频(nn.Buffer,persistent=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)),基频 base 由 rope_parameters[layer_type]["rope_theta"] 给出(同文件 L86-L96)。
4.3 注意力:QK-Norm + 滑动窗口 + 可插拔后端
MellumAttention(modeling_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.forward 在 output_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_layernorm → mlp(稀疏或稠密)→ 残差,并继承自 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_plan):q_proj/k_proj/v_proj为colwise、o_proj为rowwise;q_norm/k_norm为replicated_with_grad_allreduce;专家gate_up_proj为packed_colwise、down_proj为rowwise,且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_plan):embed_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),返回 MoeModelOutputWithPast(last_hidden_state 等) |
MellumForCausalLM |
因果语言建模头 | 额外接受 labels、output_router_logits、logits_to_keep;返回 MoeCausalLMOutputWithPast,含 loss、aux_loss、router_logits 等 |
几个参数在代码中的行为:
position_ids:未显式传入时,代码会用past_seen_tokens与序列长度自动补齐(modeling_mellum.py);logits_to_keep:仅对末尾 N 个位置计算 logits 以省显存,配合labels做自回归训练切片使用;- 初始化(
_init_weights)对MellumExperts的gate_up_proj/down_proj与MellumTopKRouter.weight单独做了均值为 0、std=initializer_range的正态初始化,并对旋转编码按层类型刷新逆频缓冲。
若要在没有官方权重时先本地跑通结构,可以直接用随机配置构建:
from transformers import MellumModel, MellumConfig
configuration = MellumConfig() # 12B 全默认参数
model = MellumModel(configuration) # 随机初始化,验证前向/结构
注意该示例会创建约 12B 参数的随机模型,仅适合小批量结构自检,勿在低资源环境下盲目执行。
八、延伸阅读路径
若要继续深入,建议在仓库内按以下顺序阅读:
- 模型文档源文件:docs/source/en/model_doc/mellum.md,对应文档目录注册见 docs/source/en/_toctree.yml;
- 手写源(改动入口):src/transformers/models/mellum/modular_mellum.py,注意其继承/组合自
qwen3_moe与laguna两个既有模块(见其头部 import),若需对比被复用的组件,可参考 src/transformers/models/qwen3_moe/ 与 src/transformers/models/laguna/; - 自动生成文件:
configuration_mellum.py与modeling_mellum.py(src/transformers/models/mellum/)为 modular 的产物,仓库 CI 会强制二者同步; - RoPE 公共工具:src/transformers/modeling_rope_utils.py(
ROPE_INIT_FUNCTIONS、RopeParameters、dynamic_rope_update的来源),理解逐层类型 RoPE 如何接入 transformers 统一的旋转编码框架。
综上,Mellum 以 Qwen3-MoE 为骨架、以“逐层类型”为编排手段,把稠密/稀疏 MLP、全量/滑动注意力与双频段 RoPE 组合进 28 层结构中,在 12B 总参数下维持 2.5B 激活的高效推理。对于代码生成、补全类场景,你可以基于上文任一示例直接替换模型标识体验;若要二次开发或微调,则需要从 modular_mellum.py 出发修改,并让改动同步到生成的建模与配置文件。
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