Skip to content
Merged
Show file tree
Hide file tree
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
8 changes: 8 additions & 0 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -1388,6 +1388,14 @@ def uses_extract_hidden_states(self) -> bool:
def use_ngram_gpu(self) -> bool:
return self.method == "ngram_gpu"

def use_multi_module_mtp(self) -> bool:
if self.method != "mtp" or self.draft_model_config is None:
return False
num_mtp_layers = getattr(
self.draft_model_config.hf_config, "num_nextn_predict_layers", 1
)
return min(num_mtp_layers, self.num_speculative_tokens) > 1

def __repr__(self) -> str:
method = self.method
model = (
Expand Down
55 changes: 42 additions & 13 deletions vllm/models/inkling/nvidia/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 @@ -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."""

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
Expand Down
25 changes: 21 additions & 4 deletions vllm/v1/worker/gpu/input_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,13 +191,16 @@ def make_dummy(
def _prepare_prefill_inputs_kernel(
input_ids_ptr,
next_prefill_tokens_ptr,
next_prefill_tokens_stride,
num_lookahead,
idx_mapping_ptr,
query_start_loc_ptr,
all_token_ids_ptr,
all_token_ids_stride,
prefill_lens_ptr,
num_computed_tokens_ptr,
BLOCK_SIZE: tl.constexpr,
LOOKAHEAD_BLOCK: tl.constexpr,
):
batch_idx = tl.program_id(0)
req_state_idx = tl.load(idx_mapping_ptr + batch_idx)
Expand All @@ -218,10 +221,20 @@ def _prepare_prefill_inputs_kernel(
tokens = tl.load(request_ptr + num_computed + block, mask=mask)
tl.store(input_ids_ptr + query_start + block, tokens, mask=mask)

next_pos = num_computed + query_len
if next_pos < prefill_len:
next_token = tl.load(request_ptr + next_pos)
tl.store(next_prefill_tokens_ptr + req_state_idx, next_token)
# Store the next num_lookahead prefill tokens.
lookahead = tl.arange(0, LOOKAHEAD_BLOCK)
pos = num_computed + query_len + lookahead
in_lookahead = lookahead < num_lookahead
tokens = tl.load(
request_ptr + pos, mask=in_lookahead & (pos < prefill_len), other=0
)
tl.store(
next_prefill_tokens_ptr
+ lookahead * next_prefill_tokens_stride
+ req_state_idx,
tokens,
mask=in_lookahead,
)


def prepare_prefill_inputs(
Expand All @@ -234,16 +247,20 @@ def prepare_prefill_inputs(
num_computed_tokens: torch.Tensor,
) -> None:
num_reqs = idx_mapping.shape[0]
num_lookahead = next_prefill_tokens.shape[0]
_prepare_prefill_inputs_kernel[(num_reqs,)](
input_ids,
next_prefill_tokens,
next_prefill_tokens.stride(0),
num_lookahead,
idx_mapping,
query_start_loc,
all_token_ids,
all_token_ids.stride(0),
prefill_len,
num_computed_tokens,
BLOCK_SIZE=1024,
LOOKAHEAD_BLOCK=triton.next_power_of_2(num_lookahead),
)


Expand Down
14 changes: 13 additions & 1 deletion vllm/v1/worker/gpu/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,15 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device):
self.is_pooling_model = self.model_config.runner_type == "pooling"
self.pooling_runner: PoolingRunner | None = None

# Multi-module MTP feeds its modules the next num_speculative_steps prefill
# tokens during chunked prefill. Other speculators only read the immediate
# next one.
num_prefill_lookahead = (
self.num_speculative_steps
if self.speculative_config is not None
and self.speculative_config.use_multi_module_mtp()
else 1
)
# General request states.
self.req_states = RequestState(
max_num_reqs=self.max_num_reqs,
Expand All @@ -231,6 +240,7 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device):
num_speculative_steps=self.num_speculative_steps,
vocab_size=self.vocab_size,
device=self.device,
num_prefill_lookahead=num_prefill_lookahead,
)
self.input_buffers = InputBuffers(
max_num_reqs=self.max_num_reqs,
Expand Down Expand Up @@ -628,7 +638,7 @@ def _dummy_run(
torch.zeros(
input_batch.num_tokens,
dtype=torch.bool,
device=self.device,
device="cpu",
Comment thread
benchislett marked this conversation as resolved.
),
)

Expand Down Expand Up @@ -1521,6 +1531,8 @@ def sample_tokens(
# NOTE: This is done here because postprocess updates
# num_computed_prefill_tokens.
# The EAGLE/MTP drafter reads one position ahead of the target.
# TODO(TheEpicDolphin): Gather MM embeddings for all speculative
# steps during multi-module MTP.
mm_inputs = self.model_state.gather_mm_embeddings(
input_batch, draft_lookahead=1
)
Expand Down
3 changes: 3 additions & 0 deletions vllm/v1/worker/gpu/spec_decode/dflash/speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from collections.abc import Mapping
from typing import Any

import numpy as np
import torch
import torch.nn as nn

Expand Down Expand Up @@ -282,6 +283,7 @@ def _build_draft_attn_metadata(
step: int,
num_query_per_req: int | None = None,
causal: bool | Mapping[int, bool] = False,
query_start_loc_np: np.ndarray | None = None,
) -> dict[str, Any] | None:
if not self.draft_attn_layer_names:
return None
Expand All @@ -294,6 +296,7 @@ def _build_draft_attn_metadata(
step=step,
num_query_per_req=self.num_query_per_req,
causal=causal,
query_start_loc_np=query_start_loc_np,
)

@torch.inference_mode()
Expand Down
2 changes: 2 additions & 0 deletions vllm/v1/worker/gpu/spec_decode/multi_module_mtp/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
Loading
Loading