Skip to content
Open
Show file tree
Hide file tree
Changes from 5 commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
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
169 changes: 169 additions & 0 deletions tests/v1/worker/test_pp_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,169 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for PPHandler broadcast/draft relay and sparse MLA stale-buffer fix."""

from types import SimpleNamespace
from unittest import mock

import pytest
import torch

from vllm.v1.worker.gpu.pp_utils import PPHandler


class FakePPHandler:
"""Minimal stand-in for PPHandler that skips CUDA/distributed setup.

Replicates the broadcast padding logic from PPHandler.broadcast().
"""

def __init__(self, max_sample_len: int, is_last_rank: bool):
self.max_sample_len = max_sample_len
self.is_last_rank = is_last_rank

def broadcast(self, sampled_token_ids: torch.Tensor) -> torch.Tensor:
assert self.is_last_rank
width = sampled_token_ids.shape[-1]
if width != self.max_sample_len:
assert width < self.max_sample_len
padded = sampled_token_ids.new_full(
(sampled_token_ids.shape[0], self.max_sample_len), -1
)
padded[:, :width] = sampled_token_ids
sampled_token_ids = padded
return sampled_token_ids


# ---------------------------------------------------------------------------
# Fix #1: DeepSeekMTP implements SupportsPP
# ---------------------------------------------------------------------------


def test_deepseek_mtp_implements_supports_pp():
"""Verify DeepSeekMTP has the SupportsPP interface (Fix #1)."""
from vllm.model_executor.models.deepseek_mtp import DeepSeekMTP
from vllm.model_executor.layers.interfaces import SupportsPP

assert issubclass(DeepSeekMTP, SupportsPP), (
"DeepSeekMTP must implement SupportsPP to be built under PP"
)


def test_deepseek_mtp_has_make_empty_intermediate_tensors():
"""Verify DeepSeekMTP provides make_empty_intermediate_tensors (Fix #1)."""
from vllm.model_executor.models.deepseek_mtp import DeepSeekMTP

assert hasattr(DeepSeekMTP, "make_empty_intermediate_tensors"), (
"DeepSeekMTP must provide make_empty_intermediate_tensors"
)


# ---------------------------------------------------------------------------
# Fix #2: PPHandler.broadcast() pads sampled_token_ids to max_sample_len
# ---------------------------------------------------------------------------


@pytest.mark.parametrize(
"width,max_sample_len",
[
(1, 2), # prefill / first decode (width=1, max_sample_len=2)
(1, 4), # K=3 spec decode (width=1, max_sample_len=4)
(2, 4), # K=1 spec decode (width=2, max_sample_len=4)
],
)
def test_pphandler_broadcast_pads_to_max_sample_len(width, max_sample_len):
"""Verify broadcast() pads sampled_token_ids to max_sample_len (Fix #2)."""
handler = FakePPHandler(max_sample_len=max_sample_len, is_last_rank=True)
num_reqs = 4
sampled = torch.zeros(num_reqs, width, dtype=torch.int64)
result = handler.broadcast(sampled)
assert result.shape == (num_reqs, max_sample_len)
# Trailing positions should be -1 (ignored by post_update)
assert (result[:, width:] == -1).all()
# Original values preserved
assert (result[:, :width] == 0).all()


def test_pphandler_broadcast_no_pad_when_already_max():
"""broadcast() should not pad when width == max_sample_len."""
handler = FakePPHandler(max_sample_len=4, is_last_rank=True)
sampled = torch.zeros(4, 4, dtype=torch.int64)
result = handler.broadcast(sampled)
assert result.shape == (4, 4)
assert (result == 0).all()


# ---------------------------------------------------------------------------
# Fix #4: Stale topk_indices_buffer is read dynamically
# ---------------------------------------------------------------------------


def test_sparse_mla_backend_reads_topk_indices_buffer_dynamically():
"""Verify sparse MLA backends read topk_indices_buffer dynamically via
self._indexer (Fix #4). After _maybe_share_lm_head replaces
Indexer.topk_indices_buffer, the impl must see the new buffer."""

# Simulate the pattern used in all sparse MLA backends:
# __init__ stores self._indexer = indexer
# forward_mqa reads self._indexer.topk_indices_buffer dynamically

initial_buffer = torch.zeros(128, 128, dtype=torch.int32)
indexer = SimpleNamespace(topk_indices_buffer=initial_buffer)

# Create a minimal impl that follows the same pattern as the real backends
class FakeSparseMLAImpl:
def __init__(self, indexer=None, topk_indices_buffer=None):
self._indexer = indexer
self.topk_indices_buffer = (
indexer.topk_indices_buffer
if indexer is not None
else topk_indices_buffer
)

def forward_mqa(self):
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
return buf

impl = FakeSparseMLAImpl(indexer=indexer)

# Simulate _maybe_share_lm_head replacing the buffer
new_buffer = torch.ones(128, 128, dtype=torch.int32)
indexer.topk_indices_buffer = new_buffer

# The impl should read the new buffer dynamically
assert impl.forward_mqa() is new_buffer, (
"Impl must read topk_indices_buffer dynamically via self._indexer"
)
assert impl.forward_mqa() is not initial_buffer, (
"Impl must not hold a stale reference to the old buffer"
)


def test_sparse_mla_backend_handles_no_indexer():
"""Verify sparse MLA backends handle indexer=None (skip-topk layers)."""

class FakeSparseMLAImpl:
def __init__(self, indexer=None, topk_indices_buffer=None):
self._indexer = indexer
self.topk_indices_buffer = (
indexer.topk_indices_buffer
if indexer is not None
else topk_indices_buffer
)

def forward_mqa(self):
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
return buf

# No indexer, buffer passed directly
buffer = torch.ones(128, 128, dtype=torch.int32)
impl = FakeSparseMLAImpl(indexer=None, topk_indices_buffer=buffer)
assert impl.forward_mqa() is buffer
17 changes: 15 additions & 2 deletions vllm/model_executor/models/deepseek_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,12 @@
_try_load_fp8_indexer_wk,
get_spec_layer_idx_from_weight_name,
)
from .utils import get_pp_missing_layer_names, maybe_prefix
from .interfaces import SupportsPP
from .utils import (
get_pp_missing_layer_names,
make_empty_intermediate_tensors_factory,
maybe_prefix,
)

logger = init_logger(__name__)

Expand Down Expand Up @@ -205,7 +210,7 @@ def compute_logits(


@support_torch_compile
class DeepSeekMTP(nn.Module, DeepseekV2MixtureOfExperts):
class DeepSeekMTP(nn.Module, DeepseekV2MixtureOfExperts, SupportsPP):
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
self.config = vllm_config.model_config.hf_config
Expand All @@ -215,6 +220,14 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
)
# Set MoE hyperparameters
self.set_moe_parameters()
# PP support: the MTP draft runs only on the last PP stage (the runner gates
# drafter construction on get_pp_group().is_last_rank), so it never actually
# consumes PP intermediate tensors — but SupportsPP requires this factory.
self.make_empty_intermediate_tensors = (
make_empty_intermediate_tensors_factory(
["hidden_states", "residual"], self.config.hidden_size
)
)

def set_moe_parameters(self):
self.num_moe_layers = self.config.num_nextn_predict_layers
Expand Down
27 changes: 25 additions & 2 deletions vllm/model_executor/models/qwen3_5_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from .interfaces import (
MultiModalEmbeddings,
SupportsMultiModal,
SupportsPP,
_require_is_multimodal,
)
from .utils import (
Expand Down Expand Up @@ -130,7 +131,8 @@ def forward(
inputs_embeds: torch.Tensor | None = None,
spec_step_idx: int = 0,
) -> torch.Tensor:
if get_pp_group().is_first_rank:
pp_group = get_pp_group()
if pp_group.is_first_rank:
if inputs_embeds is None:
inputs_embeds = self.embed_input_ids(input_ids)
assert hidden_states.shape[-1] == inputs_embeds.shape[-1]
Expand All @@ -139,7 +141,18 @@ def forward(
hidden_states = torch.cat([inputs_embeds, hidden_states], dim=-1)
hidden_states = self.fc(hidden_states)
residual = None
elif pp_group.is_last_rank:
# Last PP rank: apply the same fc projection as first rank,
# using the target model's output on this rank.
if inputs_embeds is None:
inputs_embeds = self.embed_input_ids(input_ids)
inputs_embeds = self.pre_fc_norm_embedding(inputs_embeds)
hidden_states = self.pre_fc_norm_hidden(hidden_states)
hidden_states = torch.cat([inputs_embeds, hidden_states], dim=-1)
hidden_states = self.fc(hidden_states)
residual = None
else:
# Middle PP rank: use intermediate tensors from previous rank.
assert intermediate_tensors is not None
hidden_states = intermediate_tensors["hidden_states"]
residual = intermediate_tensors["residual"]
Expand Down Expand Up @@ -354,7 +367,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
"hidden_states": 0,
}
)
class Qwen3_5MTP(LocalArgmaxMixin, nn.Module, SupportsMultiModal):
class Qwen3_5MTP(LocalArgmaxMixin, nn.Module, SupportsMultiModal, SupportsPP):
packed_modules_mapping = {
"qkv_proj": [
"q_proj",
Expand Down Expand Up @@ -397,6 +410,10 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):

self.logits_processor = LogitsProcessor(config.vocab_size)

self.make_empty_intermediate_tensors = (
self.model.make_empty_intermediate_tensors
)

def embed_input_ids(
self,
input_ids: torch.Tensor,
Expand Down Expand Up @@ -432,6 +449,12 @@ def forward(
inputs_embeds: torch.Tensor | None = None,
**kwargs: object,
):
if not get_pp_group().is_first_rank and intermediate_tensors is None:
intermediate_tensors = self.make_empty_intermediate_tensors(
batch_size=hidden_states.shape[0],
dtype=hidden_states.dtype,
device=hidden_states.device,
)
hidden_states = self.model(
input_ids, positions, hidden_states, intermediate_tensors, inputs_embeds
)
Expand Down
3 changes: 2 additions & 1 deletion vllm/models/deepseek_v32/nvidia/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,12 +142,13 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
# DSA is always sparse (has index_topk); allocate the shared top-k
# buffer the indexer writes and the sparse MLA backend reads.
self.is_v32 = True
topk_indices_buffer = torch.empty(
self.topk_indices_buffer = torch.empty(
vllm_config.scheduler_config.max_num_batched_tokens,
config.index_topk,
dtype=torch.int32,
device=self.device,
)
topk_indices_buffer = self.topk_indices_buffer

if get_pp_group().is_first_rank:
self.embed_tokens = VocabParallelEmbedding(
Expand Down
10 changes: 8 additions & 2 deletions vllm/v1/attention/backends/mla/flashattn_mla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,7 @@ def __init__(
self.kv_cache_dtype = kv_cache_dtype
self.kv_lora_rank: int = mla_args["kv_lora_rank"]
self.qk_rope_head_dim: int = mla_args["qk_rope_head_dim"]
self._indexer = indexer
self.topk_indices_buffer: torch.Tensor | None = (
indexer.topk_indices_buffer if indexer is not None else topk_indices_buffer
)
Expand All @@ -248,8 +249,13 @@ def forward_mqa(
q_nope, q_rope = q
num_actual_toks = q_rope.shape[0]

assert self.topk_indices_buffer is not None
topk_indices = self.topk_indices_buffer[:num_actual_toks]
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
assert buf is not None, "topk_indices_buffer required for sparse MLA"
topk_indices = buf[:num_actual_toks]
topk_indices, valid_counts = triton_convert_req_index_to_global_index(
attn_metadata.req_id_per_token[:num_actual_toks],
attn_metadata.block_table,
Expand Down
10 changes: 8 additions & 2 deletions vllm/v1/attention/backends/mla/flashinfer_mla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -390,6 +390,7 @@ def __init__(
self.qk_nope_head_dim: int = mla_args["qk_nope_head_dim"]
self.qk_rope_head_dim: int = mla_args["qk_rope_head_dim"]

self._indexer = indexer
# The indexer carries the shared buffer for normal layers and tests;
# the explicitly-passed buffer covers backbone skip layers, whose
# indexer is not constructed (see deepseek_v2.py).
Expand Down Expand Up @@ -418,8 +419,13 @@ def forward_mqa(

num_actual_toks = q.shape[0]

assert self.topk_indices_buffer is not None
topk_indices = self.topk_indices_buffer[:num_actual_toks]
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
assert buf is not None, "topk_indices_buffer required for sparse MLA"
topk_indices = buf[:num_actual_toks]

topk_indices_physical, seq_lens = triton_convert_req_index_to_global_index(
attn_metadata.req_id_per_token[:num_actual_toks],
Expand Down
10 changes: 8 additions & 2 deletions vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm120.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@ def __init__(
)
self.kv_scale_format = _kv_scale_format_for_model(model_type)

self._indexer = indexer
# Skip-topk layers are built with indexer=None and get the shared
# buffer via mla_args instead (cf. FLASHMLA_SPARSE).
self.topk_indices_buffer: torch.Tensor | None = (
Expand Down Expand Up @@ -112,8 +113,13 @@ def forward_mqa(

num_actual_toks = q.shape[0]

assert self.topk_indices_buffer is not None
topk_indices = self.topk_indices_buffer[:num_actual_toks]
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
assert buf is not None, "topk_indices_buffer required for sparse MLA"
topk_indices = buf[:num_actual_toks]

topk_indices_physical = cast(
torch.Tensor,
Expand Down
10 changes: 8 additions & 2 deletions vllm/v1/attention/backends/mla/flashmla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -568,6 +568,7 @@ def __init__(
self.kv_cache_dtype = kv_cache_dtype
self.kv_lora_rank: int = mla_args["kv_lora_rank"]
self.softmax_scale = scale
self._indexer = indexer
# The indexer carries the shared buffer for normal layers and tests;
# the explicitly-passed buffer covers backbone skip layers, whose
# indexer is not constructed (see deepseek_v2.py).
Expand Down Expand Up @@ -873,8 +874,13 @@ def forward_mqa(
num_actual_toks = q.shape[0]

# Get topk indices
assert self.topk_indices_buffer is not None
topk_indices = self.topk_indices_buffer[:num_actual_toks]
buf = (
self._indexer.topk_indices_buffer
if self._indexer is not None
else self.topk_indices_buffer
)
assert buf is not None, "topk_indices_buffer required for sparse MLA"
topk_indices = buf[:num_actual_toks]

use_fp8_cache = self.kv_cache_dtype == "fp8_ds_mla"

Expand Down
Loading
Loading