From 7049a5d67e65d04b38f3a1908add5068a5b4184d Mon Sep 17 00:00:00 2001 From: Andreas Karatzas Date: Mon, 3 Aug 2026 03:27:56 +0000 Subject: [PATCH] [ROCm] Restore Inkling MTP backend parity Signed-off-by: Andreas Karatzas --- vllm/models/inkling/amd/mtp.py | 55 ++++++++++++++++++++++++++-------- 1 file changed, 42 insertions(+), 13 deletions(-) diff --git a/vllm/models/inkling/amd/mtp.py b/vllm/models/inkling/amd/mtp.py index 34cde020ed38..197894beb915 100644 --- a/vllm/models/inkling/amd/mtp.py +++ b/vllm/models/inkling/amd/mtp.py @@ -2,9 +2,13 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Inkling MTP (Multi-Token Prediction) draft model (NVIDIA). -Implements the first MTP depth from the reference ``mtp_model.py`` shipped with -the checkpoint. It owns ``hidden_norm`` / ``embed_norm`` RMSNorms, a ``2H -> H`` -input projection, and a full Inkling transformer block with a dense bf16 MLP. +Mirrors the reference ``mtp_model.py`` shipped with the checkpoint: each MTP +depth ``i`` owns ``hidden_norm`` / ``embed_norm`` RMSNorms, an ``input_proj`` +(``2H -> H``) and a full Inkling transformer block (dense bf16 MLP, with the +same short convolutions as the backbone; its attention is full or sliding-window +per depth, selected by ``mtp_config.local_layer_ids``). When enabled, a shared +``chain_norm`` is applied after every depth; its output is both the logits input +and the previous hidden state fed to the next depth. The draft shares the target's token embedding table and LM head (``load_eagle_model`` wires those references) and applies the backbone @@ -53,6 +57,16 @@ def _mtp_depth_from_name(name: str) -> int | None: return int(m.group(1)) if m else None +def _select_mtp_depth_count(n_predict: int, num_spec: int | None) -> int: + num_layers = min(n_predict, num_spec) if num_spec else n_predict + if num_layers <= 0: + raise ValueError( + "Inkling MTP requires num_nextn_predict_layers and " + "num_speculative_tokens to select at least one depth layer." + ) + return num_layers + + class InklingMTPDepthLayer(nn.Module): """One MTP depth: norm both inputs, fuse (2H->H), run a Inkling block.""" @@ -96,14 +110,30 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: vllm_config.speculative_config.draft_model_config.hf_config ) self.config = config - if vllm_config.speculative_config.num_speculative_tokens != 1: - raise ValueError( - "Inkling MTP currently supports exactly one speculative token" - ) + # The checkpoint ships num_nextn_predict_layers depth blocks, but only + # the first ``num_speculative_tokens`` are exercised (step i uses depth + # i). Build only those to save memory — each depth is a full Inkling block + # with its own (large) full-history sconv caches and KV cache. + n_predict = config.num_nextn_predict_layers + num_spec = vllm_config.speculative_config.num_speculative_tokens + self.num_mtp_layers = _select_mtp_depth_count(n_predict, num_spec) self.chain_hidden_post_norm = config.chain_hidden_post_norm + + # Depth blocks whose attention is sliding-window (swa_* head config) + # rather than full; keyed by MTP depth via the checkpoint's + # mtp_config.local_layer_ids (promoted onto the draft config). Mirrors + # InklingModel's local_ids split, but over MTP depths, not backbone + # layers. local_ids = set(config.local_layer_ids) + + # Keyed by depth index (str) to mirror the checkpoint layout. self.layers = nn.ModuleDict( - {"0": InklingMTPDepthLayer(config, f"{prefix}.layers.0", 0 in local_ids)} + { + str(idx): InklingMTPDepthLayer( + config, f"{prefix}.layers.{idx}", idx in local_ids + ) + for idx in range(self.num_mtp_layers) + } ) self.chain_norm = ( InklingRMSNorm(config.hidden_size, eps=config.rms_norm_eps) @@ -201,9 +231,8 @@ def forward( # auto-enumerated as a draft attention layer); its per-token metadata is # built by the speculator's build_attn_metadata and read from the # forward context, so nothing extra is threaded here. - if spec_step_idx != 0: - raise ValueError("Inkling MTP only supports spec_step_idx=0") - layer = self.layers["0"] + depth = spec_step_idx % self.num_mtp_layers + layer = self.layers[str(depth)] combined = self.fused_input_cat( layer, previous_hidden_states, input_ids, inputs_embeds ) @@ -355,8 +384,8 @@ def _load(name: str, weight: torch.Tensor, shard_id: object = None) -> bool: # Only consume the MTP weights; everything else belongs to the target. if ".mtp." not in name: continue - # Only the first checkpoint depth is used for MTP=1. - if depth is not None and depth != 0: + # Skip depth blocks beyond the ones we built (num_speculative_tokens). + if depth is not None and depth >= module.model.num_mtp_layers: continue # model.mtp.chain_norm.weight -> model.chain_norm.weight # model.mtp.layers.{i}.X -> model.layers.{i}.X