Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 42 additions & 13 deletions vllm/models/inkling/amd/mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -54,6 +58,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."""

Expand Down Expand Up @@ -97,14 +111,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)
Expand Down Expand Up @@ -202,9 +232,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
)
Expand Down Expand Up @@ -356,8 +385,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
Expand Down
Loading