Skip to content
Closed
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
239 changes: 239 additions & 0 deletions tests/v1/spec_decode/test_llm_base_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,24 @@
``block_size``.
"""

from contextlib import nullcontext
from types import SimpleNamespace

import numpy as np
import pytest
import torch

import vllm.v1.spec_decode.llm_base_proposer as llm_base_proposer
from vllm.config import CUDAGraphMode
from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM
from vllm.v1.attention.backend import CommonAttentionMetadata
from vllm.v1.attention.backends.triton_attn import (
TritonAttentionBackend,
TritonAttentionMetadataBuilder,
)
from vllm.v1.kv_cache_interface import FullAttentionSpec
from vllm.v1.spec_decode.eagle import EagleProposer
from vllm.v1.worker.utils import AttentionGroup

SCHEDULER_BLOCK_SIZE = 256
KERNEL_BLOCK_SIZE = 64
Expand Down Expand Up @@ -110,3 +122,230 @@ def test_draft_layer_iteration_is_deterministic(monkeypatch: pytest.MonkeyPatch)
assert len(proposer.draft_attn_groups) == 1
assert proposer.draft_attn_groups[0].layer_names == expected_order
assert proposer.block_size == KERNEL_BLOCK_SIZE


@pytest.fixture
def draft_metadata():
"""Real CPU metadata builder, without loading draft weights or running kernels."""
config = SimpleNamespace(
compilation_config=SimpleNamespace(
cudagraph_mode=CUDAGraphMode.PIECEWISE, static_forward_context={}
),
model_config=SimpleNamespace(
get_num_attention_heads=lambda _: 8,
get_num_kv_heads=lambda _: 8,
get_head_size=lambda: 64,
rswa_window=None,
),
parallel_config=SimpleNamespace(
data_parallel_size=1, decode_context_parallel_size=1
),
speculative_config=None,
scheduler_config=SimpleNamespace(max_num_seqs=2),
)
spec = FullAttentionSpec(
block_size=128, num_kv_heads=8, head_size=64, dtype=torch.float32
)
names = ["draft.attn"]
builder = TritonAttentionMetadataBuilder(spec, names, config, torch.device("cpu"))
proposer = EagleProposer.__new__(EagleProposer)
proposer.method = "eagle3"
proposer.device = torch.device("cpu")
proposer.vllm_config = config
proposer.speculative_config = SimpleNamespace(disable_padded_drafter_batch=False)
proposer.parallel_drafting = False
proposer.constant_draft_positions = False
proposer.needs_extra_input_slots = False
proposer.supports_mm_inputs = False
proposer.uses_mrope = False
proposer.draft_model_config = SimpleNamespace(uses_mrope=False)
proposer.num_speculative_tokens = 3
proposer._draft_attn_layer_names = set(names)
proposer.draft_attn_groups = [
AttentionGroup(TritonAttentionBackend, names, spec, 0, [builder])
]
proposer._draft_query_start_loc_cpu_cache = {}
proposer.token_arange_np = np.arange(16, dtype=np.int32)
common = CommonAttentionMetadata(
query_start_loc=torch.tensor([0, 2, 5], dtype=torch.int32),
query_start_loc_cpu=torch.tensor([0, 2, 5], dtype=torch.int32),
seq_lens=torch.tensor([10, 20], dtype=torch.int32),
num_reqs=2,
num_actual_tokens=5,
max_query_len=3,
max_seq_len=20,
block_table_tensor=torch.tensor([[0, 1], [2, 3]], dtype=torch.int32),
slot_mapping=torch.tensor([9, 10, 275, 276, 277], dtype=torch.int64),
)
return proposer, builder, common


def _propose_until_decode(monkeypatch, proposer, builder, common, k=3):
"""Run propose through real metadata construction; omit GPU/model execution."""

class DraftDecodeReady(Exception):
pass

class TestEagle(Eagle3LlamaForCausalLM):
def __init__(self):
torch.nn.Module.__init__(self)

def combine_hidden_states(self, hidden_states):
return hidden_states

def forward(self, **kwargs):
hidden = torch.zeros(common.num_actual_tokens, 2)
return hidden, hidden

proposer.model = TestEagle()
proposer.hidden_size = 2
proposer._share_mtp_indices = False
proposer.eplb_state = None
proposer.positions = torch.arange(16)
proposer.arange = torch.arange(16, dtype=torch.int32)
proposer.allowed_attn_types = None
proposer.block_size = 128
proposer.use_heterogeneous_vocab = False
monkeypatch.setattr(
llm_base_proposer, "set_forward_context", lambda *a, **kw: nullcontext()
)
monkeypatch.setattr(proposer, "_get_slot_mapping", lambda *a: None)
monkeypatch.setattr(
proposer,
"_determine_batch_execution_and_padding",
lambda n: (CUDAGraphMode.PIECEWISE, n, None),
)
monkeypatch.setattr(
proposer,
"set_inputs_first_pass",
lambda **kw: (common.num_actual_tokens, torch.arange(common.num_reqs), common),
)
monkeypatch.setattr(
proposer,
"build_model_inputs_first_pass",
lambda *a: ({}, common.num_actual_tokens),
)
monkeypatch.setattr(
proposer,
"_sample_draft_tokens",
lambda *a: (torch.zeros(common.num_reqs, dtype=torch.int64), None),
)

def update_positions(positions, metadata, *args):
metadata.seq_lens.add_(1)
metadata.max_seq_len += 1
return positions + 1

monkeypatch.setattr(
proposer, "_update_positions_dependent_metadata", update_positions
)

build = builder.build_for_drafting
captured = []

def capture(common_attn_metadata, draft_index):
result = build(common_attn_metadata, draft_index)
if draft_index == 1:
captured.append(result)
raise DraftDecodeReady
return result

with monkeypatch.context() as patch:
patch.setattr(builder, "build_for_drafting", capture)
with pytest.raises(DraftDecodeReady):
proposer.propose(
k,
torch.zeros(common.num_reqs, dtype=torch.int32),
torch.arange(common.num_reqs),
torch.zeros(common.num_reqs, 2),
torch.zeros(common.num_reqs, dtype=torch.int32),
None,
common,
None,
)
metadata = captured[0]
assert metadata.max_query_len == 1
assert metadata.num_actual_tokens == common.num_reqs
assert metadata.seq_lens is common.seq_lens
assert metadata.max_seq_len == common.max_seq_len
assert metadata.block_table is common.block_table_tensor
assert metadata.slot_mapping is common.slot_mapping
assert common.query_start_loc_cpu.tolist() == list(range(common.num_reqs + 1))
return common.query_start_loc_cpu


def test_propose_reuses_cpu_query_offsets_across_batch_and_draft_length_changes(
monkeypatch, draft_metadata
):
proposer, builder, common = draft_metadata
first = _propose_until_decode(monkeypatch, proposer, builder, common)
smaller = common.replace(
num_reqs=1, num_actual_tokens=1, seq_lens=common.seq_lens[:1]
)
_propose_until_decode(monkeypatch, proposer, builder, smaller, k=2)
common.seq_lens = torch.tensor([40, 50], dtype=torch.int32)
common.max_seq_len = 50
common.slot_mapping = common.slot_mapping + 2
second = _propose_until_decode(monkeypatch, proposer, builder, common, k=4)
assert second is first
assert common.seq_lens.tolist() == [41, 51]
assert common.max_seq_len == 51
proposer.token_arange_np[:] = -1
assert first.tolist() == [0, 1, 2]


def test_propose_cpu_offsets_fall_back_for_unrecognized_builder(
monkeypatch, draft_metadata
):
proposer, builder, common = draft_metadata
cached = _propose_until_decode(monkeypatch, proposer, builder, common)

class OtherBuilder(TritonAttentionMetadataBuilder):
pass

builder.__class__ = OtherBuilder
first = _propose_until_decode(monkeypatch, proposer, builder, common)
first.fill_(-1)
second = _propose_until_decode(monkeypatch, proposer, builder, common)
assert second is not first
assert second.tolist() == [0, 1, 2]
assert cached.tolist() == [0, 1, 2]


@pytest.mark.parametrize("num_dims", [3, 4])
def test_cpu_query_offsets_mrope_fallback_has_private_storage(draft_metadata, num_dims):
proposer, _, _ = draft_metadata
proposer.draft_model_config = SimpleNamespace(
uses_mrope=True, mrope_num_dims=num_dims
)
proposer.uses_mrope = proposer.draft_model_config.uses_mrope
first = proposer._get_draft_query_start_loc_cpu(2)
first.fill_(-1)
assert proposer._get_draft_query_start_loc_cpu(2).tolist() == [0, 1, 2]


def test_cpu_query_offsets_released_on_backend_reinitialization(
monkeypatch, draft_metadata
):
proposer, builder, _ = draft_metadata
first = proposer._get_draft_query_start_loc_cpu(2)
monkeypatch.setattr(
llm_base_proposer,
"get_layers_from_vllm_config",
lambda *a, **kw: {
"draft.attn": SimpleNamespace(
get_attn_backend=lambda: TritonAttentionBackend
)
},
)
kv_config = SimpleNamespace(
kv_cache_groups=[
SimpleNamespace(
layer_names=["draft.attn"], kv_cache_spec=builder.kv_cache_spec
)
]
)
proposer.initialize_attn_backend(kv_config)
second = proposer._get_draft_query_start_loc_cpu(2)
assert second is not first
assert second.tolist() == [0, 1, 2]
34 changes: 30 additions & 4 deletions vllm/v1/spec_decode/llm_base_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,11 @@
from vllm.utils.torch_utils import PIN_MEMORY, async_tensor_h2d
from vllm.v1.attention.backend import CommonAttentionMetadata
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.v1.attention.backends.triton_attn import TritonAttentionMetadata
from vllm.v1.attention.backends.triton_attn import (
TritonAttentionBackend,
TritonAttentionMetadata,
TritonAttentionMetadataBuilder,
)
from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher
from vllm.v1.kv_cache_interface import KVCacheConfig, UniformTypeKVCacheSpecs
from vllm.v1.sample.metadata import SamplingMetadata
Expand Down Expand Up @@ -153,6 +157,7 @@ def __init__(
)

self.draft_attn_groups: list[AttentionGroup] = []
self._draft_query_start_loc_cpu_cache: dict[int, torch.Tensor] = {}
self.kv_cache_gid: int = -1
self.eagle3_use_aux_hidden_state: bool = (
self._get_eagle3_use_aux_hidden_state_from_config()
Expand Down Expand Up @@ -659,9 +664,9 @@ def propose(
common_attn_metadata.num_actual_tokens = batch_size
common_attn_metadata.max_query_len = 1
common_attn_metadata.query_start_loc = self.arange[: batch_size + 1]
common_attn_metadata.query_start_loc_cpu = torch.from_numpy(
self.token_arange_np[: batch_size + 1]
).clone()
common_attn_metadata.query_start_loc_cpu = self._get_draft_query_start_loc_cpu(
batch_size
)

# In padded drafter batch, we need to adjust the sequence lengths
# to remove the "padding" (i.e. rejected tokens).
Expand Down Expand Up @@ -991,6 +996,26 @@ def build_per_group_and_layer_attn_metadata(
per_layer_attn_metadata[layer_name] = attn_metadata
return per_group_attn_metadata, per_layer_attn_metadata

def _get_draft_query_start_loc_cpu(self, batch_size: int) -> torch.Tensor:
# Triton does not mutate CPU query offsets. Other builders keep receiving
# private storage, including extensions with their own metadata handling.
can_reuse = (
self.method == "eagle3"
and not self.uses_mrope
and len(self.draft_attn_groups) == 1
and self.draft_attn_groups[0].backend is TritonAttentionBackend
and type(self.draft_attn_groups[0].get_metadata_builder())
is TritonAttentionMetadataBuilder
)
if can_reuse:
cached = self._draft_query_start_loc_cpu_cache.get(batch_size)
if cached is not None:
return cached
offsets = torch.from_numpy(self.token_arange_np[: batch_size + 1]).clone()
if can_reuse:
self._draft_query_start_loc_cpu_cache[batch_size] = offsets
return offsets

def model_returns_tuple(self) -> bool:
if self.method == "mtp":
# These models return separate hidden states for logits and for
Expand Down Expand Up @@ -1721,6 +1746,7 @@ def initialize_attn_backend(
Initialize AttentionGroups for draft layers using kv_cache_config.
Called from the model runner's initialize_metadata_builders.
"""
self._draft_query_start_loc_cpu_cache = {}
all_attn_layers = get_layers_from_vllm_config(
self.vllm_config,
AttentionLayerBase, # type: ignore[type-abstract]
Expand Down
Loading