Prhub

#48869 [Model] Add Inkling MTP=1 support [3/N]

原始 PR 作者 WoosukKwon 合并时间 2026-07-17 04:27 文件变更 11 提交数 1 评论 0 代码增减 +710 / -6

执行摘要

为 Inkling 模型添加 MTP=1 推测解码支持

Inkling 模型需要 speculative decoding 提升推理吞吐。本 PR 是 Inkling MTP 支持系列的第三个切片(继 #48799 基础模型和 #48822 CUDA 图之后),添加 MTP draft 层,使模型能利用一个推测 token 加速解码。

建议精读此 PR,了解 vLLM 如何为新模型集成 speculative decoding 支持,尤其是配置覆写、权重共享和融合核函数的设计模式。

讨论亮点

PR 提交后无实质性 Review 讨论,仅 claude[bot] 自动回复。所有 CI 和验证均通过。

实现拆解

  1. 配置集成:在 speculative.py 中添加 inkling_mtp 类型,并在 hf_config_override 中解析 checkpoint 的 mtp_config,设置 n_predict=1 和相关配置。
  2. 核心模块:新增 mtp.py,定义 InklingMTPDepthLayer(单层 MTP 深度,含双 RMSNorm、输入投影和强制密集 MLP 的 Inkling 解码层)和 InklingMultiTokenPredictor(预测器,持有该层并共享 target 的 embedding 表和 LM head)。
  3. 融合算子:在 norm.py 中新增 Triton 核函数 _embed_dual_rmsnorm_cat_kernel,实现 fused gather + dual RMSNorm + concat,减少 launch 开销。
  4. 模型适配:修改 model.py,为 InklingDecoderLayer 添加 force_dense_mlp(MTP 层强制密集 MLP)和 defer_mlp_add(延迟残差添加)参数。
  5. 注册与导出:在 init.py 和 registry.py 中注册 InklingMTP 架构。
  6. 测试配套:新增 test_mtp_input_fusion.py(32 个 bit-exact 测试),补充配置覆写、注册表和契约测试。
文件 模块 状态 重要度
vllm/models/inkling/nvidia/mtp.py MTP 层 added 9.36
tests/models/inkling/test_mtp_input_fusion.py 测试 added 7.72
vllm/config/speculative.py 配置 modified 6.67
vllm/models/inkling/nvidia/ops/norm.py 融合算子 modified 6.55
vllm/models/inkling/nvidia/model.py 模型层 modified 6.49
vllm/models/inkling/__init__.py 模块入口 modified 5.94
vllm/models/inkling/configs.py 配置 modified 5.94
tests/config/test_speculative_draft_hf_overrides.py 测试 modified 5.67
vllm/model_executor/models/registry.py 注册表 modified 4.93
tests/models/inkling/test_contract_validation.py 测试 modified 4.83
tests/models/registry.py 测试 modified 4.62

关键符号

_mtp_depth_from_name InklingMTPDepthLayer.__init__ InklingMTPDepthLayer.forward InklingMultiTokenPredictor.__init__ InklingMultiTokenPredictor.embed_input_ids InklingMultiTokenPredictor.fused_input_cat InklingMTP embed_dual_rmsnorm_cat _embed_dual_rmsnorm_cat_kernel embed_rmsnorm test_embed_dual_rmsnorm_cat test_embed_rmsnorm test_inkling_override_exposes_only_first_mtp_depth test_inkling_mtp_chain_norm_is_disabled_by_default

关键源码片段

vllm/models/inkling/nvidia/mtp.py data-contract

新增 Inkling MTP draft 模型的核心实现,定义 InklingMTPDepthLayer 和 InklingMultiTokenPredictor,实现融合输入构造和权重加载。

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Inkling MTP (Multi-Token Prediction) draft model (NVIDIA)."""from __future__ import annotationsimport regex as re
import torch
from torch import nnfrom vllm.config import VllmConfig
from vllm.model_executor.layers.linear import ReplicatedLinear
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
from vllm.model_executor.models.utils import maybe_prefix
from vllm.sequence import IntermediateTensorsfrom ..configs import InklingModelConfig
from .layernorm import InklingRMSNorm
from .model import InklingDecoderLayer, InklingReplicatedEmbedding
from .ops.norm import embed_dual_rmsnorm_cat, embed_rmsnorm_ATTENTION_PARAMS_MAPPING = [
    ("qkvr", "wq_du", 0),
    ("qkvr", "wk_dv", 1),
    ("qkvr", "wv_dv", 2),
    ("qkvr", "wr_du", 3),
]def _mtp_depth_from_name(name: str) -> int | None:
    m = re.search(r"\.mtp\.layers\.(\d+)\.", name)
    return int(m.group(1)) if m else Noneclass InklingMTPDepthLayer(nn.Module):
    """One MTP depth: norm both inputs, fuse (2H->H), run a Inkling block."""
    def __init__(self, config: InklingModelConfig, prefix: str, is_local: bool) -> None:
        super().__init__()
        self.hidden_norm = InklingRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.embed_norm = InklingRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.input_proj = ReplicatedLinear(
            config.hidden_size * 2,
            config.hidden_size,
            bias=False,
            return_bias=False,
            prefix=f"{prefix}.input_proj",
        )
        self.transformer_block = InklingDecoderLayer(
            config,
            layer_id=0,
            is_local=is_local,
            quant_config=None,
            prefix=f"{prefix}.transformer_block",
            nvfp4_config=None,
            force_dense_mlp=True,
        )
    def forward(self, combined: torch.Tensor, positions: torch.Tensor) -> torch.Tensor:
        hidden = self.input_proj(combined)
        return self.transformer_block(positions, hidden)
vllm/models/inkling/nvidia/ops/norm.py infrastructure

新增 Triton 融合核函数 embed_dual_rmsnorm_cat,将两个 RMSNorm 和一个 concat 合并为单个 kernel,减少 launch 开销。

import torch
import triton
import triton.language as tl@triton.jit
def _embed_dual_rmsnorm_cat_kernel(
    hidden_ptr, # [T, N] or [T, N] (GATHER: [V, N])
    emb_ptr, # [T, N] or [V, N]
    ids_ptr, # [T] token ids (GATHER only)
    w_hidden_ptr, # [N]
    w_pre_ptr, # [N] chain pre-norm (HAS_PRE_NORM only)
    w_embed_ptr, # [N]
    out_ptr, # [T, 2N]: [rmsnorm(hidden) | rmsnorm(rmsnorm?(emb))]
    eps,
    hidden_stride_0,
    emb_stride_0,
    n_cols,
    block_size_n: tl.constexpr,
    GATHER: tl.constexpr,
    HAS_PRE_NORM: tl.constexpr,
):
    """Triton kernel: compute rmsnorm for hidden (left) and embedding (right) in one launch."""
    pid_m = tl.program_id(0).to(tl.int64)
    which = tl.program_id(1) # 0 => hidden; 1 => emb
    offs_n = tl.arange(0, block_size_n)
    mask_n = offs_n < n_cols
    if which == 0:
        x = tl.load(hidden_ptr + pid_m * hidden_stride_0 + offs_n, mask=mask_n, other=0.0).to(tl.float32)
        w = tl.load(w_hidden_ptr + offs_n, mask=mask_n, other=0.0).to(tl.float32)
    else:
        row = tl.load(ids_ptr + pid_m).to(tl.int64) if GATHER else pid_m
        x = tl.load(emb_ptr + row * emb_stride_0 + offs_n, mask=mask_n, other=0.0).to(tl.float32)
        if HAS_PRE_NORM:
            w_pre = tl.load(w_pre_ptr + offs_n, mask=mask_n, other=0.0).to(tl.float32)
            rstd = tl.math.rsqrt(tl.sum(x * x, axis=0) / n_cols + eps)
            x = (x * rstd * w_pre).to(out_ptr.dtype.element_ty).to(tl.float32)
        w = tl.load(w_embed_ptr + offs_n, mask=mask_n, other=0.0).to(tl.float32)
    rstd = tl.math.rsqrt(tl.sum(x * x, axis=0) / n_cols + eps)
    tl.store(
        out_ptr + pid_m * (2 * n_cols) + which * n_cols + offs_n,
        (x * rstd * w).to(out_ptr.dtype.element_ty),
        mask=mask_n,
    )def embed_dual_rmsnorm_cat(
    hidden: torch.Tensor,
    hidden_weight: torch.Tensor,
    embed_weight: torch.Tensor,
    eps: float,
    *,
    embeds: torch.Tensor | None = None,
    input_ids: torch.Tensor | None = None,
    embed_table: torch.Tensor | None = None,
    pre_norm_weight: torch.Tensor | None = None,
) -> torch.Tensor:
    """Fused cat([rmsnorm(hidden, w_h), rmsnorm(pre?(emb), w_e)], dim=-1)."""
    T, n = hidden.shape
    if embeds is not None:
        src, ids, src_stride = embeds, embeds, embeds.stride(0)
        gather = False
    else:
        assert input_ids is not None and embed_table is not None
        src, ids, src_stride = embed_table, input_ids, embed_table.stride(0)
        gather = True
    out = torch.empty((T, 2 * n), dtype=hidden.dtype, device=hidden.device)
    grid = (T, 2)
    _embed_dual_rmsnorm_cat_kernel[grid](
        hidden, src, ids, hidden_weight, pre_norm_weight, embed_weight,
        out, eps, hidden.stride(0), src_stride, n,
        block_size_n=triton.next_power_of_2(n),
        GATHER=gather, HAS_PRE_NORM=pre_norm_weight is not None,
    )
    return out

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

仅支持 MTP=1,若配置 num_speculative_tokens>1 会抛出明确异常;融合核函数依赖 Triton 和 CUDA,非 GPU 环境不可用;缺少大规模压力测试;权重共享引用目标模型表,避免了额外显存开销。

影响限于 Inkling 模型用户;启用推测解码后可提升推理吞吐(GSM8K 测试中 MTP 接受率约 83-84%);团队需维护新增的 MTP 模型代码和融合算子。

仅支持 MTP=1 依赖 Triton/CUDA 硬件 缺少大规模压力测试 融合核函数 bit-exact 性依赖测试验证

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论