Prhub

#27375 [Model] Add support for JetBrains' Mellum v2 code generation model

原始 PR 作者 shadeMe 合并时间 2026-07-14 13:54 文件变更 7 提交数 11 评论 11 代码增减 +706 / -8

执行摘要

新增 JetBrains Mellum v2 代码生成模型推理支持

添加对 JetBrains/Mellum2-12B-A2.5B-Thinking 及相关变体的支持,该模型采用混合滑动窗口注意力和 MoE 架构,专为代码生成任务设计。PR 从 HuggingFace 加载模型并利用 SGLang 推理框架提供服务。

值得关注本 PR 如何通过继承和配置别名降低新模型集成成本,对于类似基于 QwenMoE 架构的模型添加有参考价值。同时 review 中的讨论展示了代码兼容性和正确性审查的必要性。建议后续补充更多端到端测试并考虑 NPU 支持。

讨论亮点

Review 中主要讨论集中在三个问题:

  • gate_up_proj 应定义在类级 packed_modules_mapping 而非实例级,以避免量化框架和权重加载器出错。最终修复。
  • _MellumConfigAlias.__post_init__ 签名兼容性:初始版本使用 **kwargs 引发 TypeError,尝试修复后又导致 head_dim 回归,最终版本通过调整参数避免问题。
  • 代码简化和位置提升建议:get_hybrid_layer_ids 分支合并已完成,位置提升作为 TODO 保留。

实现拆解

  1. 模型核心实现 (python/sglang/srt/models/mellum.py): 新建 MellumForCausalLM 类,继承自 Qwen3MoeForCausalLM。子类 MellumAttention 重写 init 和 forward_prepare_npu,注入逐层滑动窗口和 per-layer-type RoPE;MellumMLP 修复 forward 签名兼容性。

  2. 配置兼容层 (python/sglang/srt/utils/hf_transformers/common.py): 在 _CONFIG_REGISTRY 中注册 'mellum' 键。优先使用 transformers>=5.10.2 的原生 MellumConfig,否则创建基于 Qwen3MoeConfig 的别名并重写 post_init 以确保 sliding_window 不会被错误清零。

  3. 模型配置路由 (python/sglang/srt/configs/model_config.py): 将 MellumForCausalLM 加入混合 SWA 模型集合并统一化 get_hybrid_layer_ids 的分支逻辑,消除了之前为 LagunaForCausalLM 单独编写的冗余代码。

  4. MoE Bench 集成 (benchmark/kernels/fused_moe_triton/common_utils.py): 在模型专家数查询表中添加 MellumForCausalLM,使其能正确参与 MoE 性能基准测试。

  5. 测试与文档 (test/registered/unit/configs/test_model_config.py, test/registered/unit/models/test_mellum.py, docs_new/docs/supported-models/generative_models.mdx): 添加针对混合层 ID 解析和位置准备的单元测试;文档中新增 Mellum 模型支持行。

文件 模块 状态 重要度
python/sglang/srt/models/mellum.py 模型定义 added 9.36
python/sglang/srt/utils/hf_transformers/common.py 配置兼容 modified 7.32
python/sglang/srt/configs/model_config.py 配置路由 modified 6.22
test/registered/unit/configs/test_model_config.py 配置测试 added 6.88
test/registered/unit/models/test_mellum.py 模型测试 added 6.8
benchmark/kernels/fused_moe_triton/common_utils.py 基准配置 modified 4.3
docs_new/docs/supported-models/generative_models.mdx 文档 modified 2.99

关键符号

MellumMLP.forward MellumAttention.__init__ MellumAttention.forward_prepare_npu _get_rope_type _compute_yarn_from_rope_params get_attention_sliding_window_size _MellumConfigAlias.__post_init__ get_hybrid_layer_ids

关键源码片段

python/sglang/srt/models/mellum.py data-contract

模型核心实现,新增 MellumForCausalLM 等类,定义混合注意力逻辑。

# MellumMLP: 简单的 forward 签名适配层
# Qwen3MoeDecoderLayer.forward 调用 self.mlp(x, forward_batch),
# 但 Qwen3MoeMLP.forward 只接受 x。
# MellumMLP 桥接这一差异以兼容父类 DecoderLayer。
class MellumMLP(Qwen3MoeMLP):
    def forward(self, x, forward_batch=None):
        return super().forward(x)
​
​
def _get_rope_type(rope_params: Dict[str, Any]) -> str:
    '''提取 rope_type,优先使用 'rope_type' 键,回退到 'type' 或默认值。'''
    return rope_params.get('rope_type') or rope_params.get('type') or 'default'
python/sglang/srt/utils/hf_transformers/common.py dependency-wiring

配置兼容层,处理 MellumConfig 注册,确保旧版 transformers 下也能加载模型。

# 当 transformers 版本过低时,使用基于 Qwen3MoeConfig 的别名类
# 并保留 sliding_window 属性以免被父类 __post_init__ 清除
if _HFMellumConfig is not None:
    _CONFIG_REGISTRY['mellum'] = _HFMellumConfig
else:
    from transformers import Qwen3MoeConfig as _HFQwen3MoeConfig
​
    class _MellumConfigAlias(_HFQwen3MoeConfig):
        model_type = 'mellum'
​
        def __post_init__(self, **kwargs):
            # Qwen3MoeConfig.__post_init__ 会清除 sliding_window,
            # 除非 use_sliding_window=True。Mellum 通过 layer_types
            # 逐层控制滑动注意力,因此必须保留 sliding_window。
            sliding_window = getattr(self, 'sliding_window', None)
            super().__post_init__(**kwargs)
            self.sliding_window = sliding_window
​
    _CONFIG_REGISTRY['mellum'] = _MellumConfigAlias

评论区精华

gate_up_proj 定义位置 设计

gemini-code-assist[bot] 指出 gate_up_proj 应在类级 packed_modules_mapping 中定义,避免量化框架和权重加载器出错。

结论:Jiminator 在后续提交中将 gate_up_proj 移至类级别 packed_modules_mapping,问题解决。 · 已解决

__post_init__ 签名兼容 正确性

gemini-code-assist[bot] 指出 __post_init__ 不应接受 **kwargs 以避免 TypeError。shadeMe 尝试修复但引发了 head_dim 参数的回归,最终 revert。

结论:最终版本保留了 __post_init__ 但调整了参数调用方式,避免了 **kwargs 导致的错误。 · 已解决

代码简化与位置提升 设计

alexnails 建议将 get_hybrid_layer_ids 中 Mellum 的分支与 Gemma4 等合并,以及考虑将位置计算提前到模型层面。

结论:分支合并已完成(统一使用 layer_types 逻辑),位置提前作为 TODO 保留。 · 已解决

风险与影响

配置兼容风险:旧版 transformers 回退路径中的 __post_init__ 签名问题曾引发回归,虽已修正但需谨慎验证。NPU 路径被显式拒绝,在 NPU 设备上直接使用会出错。单元测试仅覆盖位置预备和混合层 ID 解析,未包含注意力计算和 MLP 端到端正确性测试。新模型注册可能被未知 transformers 版本打断。

对用户:支持使用 Mellum v2 模型进行推理,调用方式与 Qwen3MoE 模型一致。对系统:新增模型配置注册影响可忽略。对团队:维护了代码一致性,但新增模型特定代码需要持续维护。对已有功能:完全向后兼容。

配置兼容风险 NPU 未测试 测试覆盖有限 依赖旧版 transformers 回退

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论