Files
enginex-ascend-910-vllm/tests/ut/spec_decode/a2/test_eagle_proposer.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

4559 lines
213 KiB
Python

# ruff: noqa: E501
import inspect
import unittest
from dataclasses import dataclass
from typing import Any
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
import torch
from vllm.config import CacheConfig, CompilationMode, CUDAGraphMode, VllmConfig, set_current_vllm_config
from vllm.forward_context import BatchDescriptor
from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM
from vllm.platforms import current_platform
from vllm.v1.spec_decode.draft_model import DraftModelProposer
import vllm_ascend.spec_decode.llm_base_proposer as llm_base_proposer
from tests.ut.base import TestBase
from vllm_ascend.ascend_config import clear_ascend_config, init_ascend_config
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
from vllm_ascend.spec_decode.draft_proposer import AscendDraftModelProposer
from vllm_ascend.spec_decode.eagle_proposer import AscendEagleProposer
from vllm_ascend.utils import enable_custom_op
from vllm_ascend.worker.pcp_utils import PCPManager, PCPSpecDecodeFirstPassInputs
enable_custom_op()
# vLLM #40732 moved `SpecDecodeBaseProposer` (and its `CpuGpuBuffer` import)
# out of `vllm.v1.spec_decode.eagle` into `vllm.v1.spec_decode.llm_base_proposer`.
_CPU_GPU_BUFFER_TARGET = "vllm.v1.spec_decode.llm_base_proposer.CpuGpuBuffer"
BLOCK_SIZE = 16
@dataclass
class BatchSpec:
"""Specification for a batch configuration (workload shape only)."""
seq_lens: list[int]
query_lens: list[int]
name: str = "unnamed"
@property
def batch_size(self):
return len(self.seq_lens)
def __post_init__(self):
assert len(self.seq_lens) == len(self.query_lens)
def compute_num_tokens(self):
return sum(self.query_lens)
def create_common_attn_metadata(
batch_spec: BatchSpec,
block_size: int,
device: torch.device,
max_block_idx: int = 1000,
arange_block_indices: bool = False,
) -> AscendCommonAttentionMetadata:
"""Create AscendCommonAttentionMetadata from a BatchSpec and ModelParams."""
# Create query start locations
query_start_loc = torch.zeros(batch_spec.batch_size + 1, dtype=torch.int32, device=device)
query_start_loc[1:] = torch.tensor(batch_spec.query_lens, dtype=torch.int32, device=device).cumsum(0)
query_start_loc_cpu = query_start_loc.cpu()
num_tokens = batch_spec.compute_num_tokens()
# Create sequence lengths
seq_lens = torch.tensor(batch_spec.seq_lens, dtype=torch.int32, device=device)
seq_lens_cpu = seq_lens.cpu()
max_seq_len = int(seq_lens_cpu.max())
# Create computed tokens (context length for each sequence)
context_lens = [batch_spec.seq_lens[i] - batch_spec.query_lens[i] for i in range(batch_spec.batch_size)]
num_computed_tokens_cpu = torch.tensor(context_lens, dtype=torch.int32)
# Create block table and slot mapping
max_blocks = (max(batch_spec.seq_lens) + block_size - 1) // block_size
if arange_block_indices:
num_blocks = batch_spec.batch_size * max_blocks
block_table_tensor = torch.arange(num_blocks, dtype=torch.int32, device=device).view(
batch_spec.batch_size, max_blocks
)
slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device).view(num_tokens)
else:
block_table_tensor = torch.randint(
0,
max_block_idx,
(batch_spec.batch_size, max_blocks),
dtype=torch.int32,
device=device,
)
slot_mapping = torch.randint(0, max_block_idx, (num_tokens,), dtype=torch.int64, device=device)
# Calculate max query length
max_query_len = max(batch_spec.query_lens)
# Create positions tensor (position indices for each token)
positions_list: list[int] = []
for i in range(batch_spec.batch_size):
seq_len = batch_spec.seq_lens[i]
query_len = batch_spec.query_lens[i]
start_pos = seq_len - query_len
positions_list.extend(range(start_pos, seq_len))
positions = torch.tensor(positions_list, dtype=torch.int64, device=device)
return AscendCommonAttentionMetadata(
query_start_loc=query_start_loc,
query_start_loc_cpu=query_start_loc_cpu,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
num_computed_tokens_cpu=num_computed_tokens_cpu,
num_reqs=batch_spec.batch_size,
num_actual_tokens=num_tokens,
max_query_len=max_query_len,
max_seq_len=max_seq_len,
block_table_tensor=block_table_tensor,
slot_mapping=slot_mapping,
causal=True,
positions=positions,
)
def assert_attr_equal(attr: str | tuple[str, Any, Any], expect: Any, actual: Any) -> None:
if isinstance(attr, tuple):
attr_name, expect_indices, actual_indices = attr
else:
attr_name = attr
expect_indices = actual_indices = None
expect_value = getattr(expect, attr_name) if expect_indices is None else getattr(expect, attr_name)[expect_indices]
actual_value = getattr(actual, attr_name) if actual_indices is None else getattr(actual, attr_name)[actual_indices]
if torch.is_tensor(expect_value) and torch.is_tensor(actual_value):
assert expect_value.device == actual_value.device, f"{attr_name} tensor device mismatch"
assert torch.equal(expect_value, actual_value), f"{attr_name} tensor mismatch"
else:
assert expect_value == actual_value, f"{attr_name} value mismatch"
def test_prepare_inputs_padded_preserves_internal_seq_lens_cpu():
proposer = AscendEagleProposer.__new__(AscendEagleProposer)
proposer.pcp_size = 1
proposer.arange = torch.arange(16, dtype=torch.int32)
proposer.runner = MagicMock()
proposer.runner.pcp_manager = None
proposer.runner.actual_seq_lengths_q = [3, 3]
proposer.runner.attn_state = AscendAttentionState.SpecDecoding
proposer.runner.decode_token_per_req = 4
internal_seq_lens_cpu = torch.tensor([7, 9], dtype=torch.int32)
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=torch.tensor([0, 3, 6], dtype=torch.int32),
query_start_loc_cpu=torch.tensor([0, 3, 6], dtype=torch.int32),
seq_lens=torch.tensor([7, 9], dtype=torch.int32),
_seq_lens_cpu=internal_seq_lens_cpu,
seq_lens_cpu=None,
num_computed_tokens_cpu=None,
num_reqs=2,
num_actual_tokens=6,
num_input_tokens=6,
max_query_len=3,
actual_seq_lengths_q=[3, 3],
block_table_tensor=torch.zeros((2, 1), dtype=torch.int32),
slot_mapping=torch.arange(6, dtype=torch.int32),
positions=torch.arange(6),
attn_state=AscendAttentionState.SpecDecoding,
decode_token_per_req=4,
max_seq_len=9,
)
spec_decode_metadata = MagicMock()
spec_decode_metadata.cu_num_draft_tokens = torch.tensor([2, 3], dtype=torch.int32)
valid_sampled_tokens_count = torch.tensor([3, 1], dtype=torch.int32)
with patch.object(llm_base_proposer, "HAS_TRITON", False):
spec_common_attn_metadata, *_ = proposer.prepare_inputs_padded(
common_attn_metadata,
spec_decode_metadata,
valid_sampled_tokens_count,
)
assert spec_common_attn_metadata._seq_lens_cpu is internal_seq_lens_cpu
assert spec_common_attn_metadata.seq_lens_cpu is None
class TestEagleProposerInitialization(TestBase):
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.cache_config = MagicMock(spec=CacheConfig)
self.vllm_config.scheduler_config = MagicMock()
self.vllm_config.model_config = MagicMock()
self.vllm_config.model_config.hf_text_config = MagicMock(
spec=[]
) # Empty spec to prevent hasattr from returning True
self.vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
self.vllm_config.compilation_config = MagicMock()
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.num_speculative_tokens = 2
self.vllm_config.speculative_config.parallel_drafting = False
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(2)])
self.vllm_config.speculative_config.draft_model_config.hf_config = MagicMock(spec=[])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Set the current vllm config
set_current_vllm_config(self.vllm_config)
def tearDown(self):
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
# Clear the current vllm config
set_current_vllm_config(None)
def test_initialization_eagle_graph(self):
self.vllm_config.speculative_config.method = "eagle"
self.vllm_config.speculative_config.draft_model_config.get_hidden_size.return_value = 4096
self.vllm_config.speculative_config.draft_model_config.get_inputs_embeds_size.return_value = 4096
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.compilation_config.mode = CompilationMode.VLLM_COMPILE
self.vllm_config.model_config.enforce_eager = False
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.speculative_config.enforce_eager = False
self.vllm_config.scheduler_config.async_scheduling = False
init_ascend_config(self.vllm_config)
with set_current_vllm_config(self.vllm_config):
proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.assertEqual(proposer.hidden_size, 4096)
self.assertTrue(proposer.use_cuda_graph)
expected_max_num_tokens = proposer.max_num_tokens
self.assertEqual(proposer.input_ids.shape, (expected_max_num_tokens,))
self.assertEqual(proposer.positions.shape, (expected_max_num_tokens,))
self.assertEqual(proposer.hidden_states.shape, (expected_max_num_tokens, 4096))
self.assertEqual(proposer.arange.shape, (expected_max_num_tokens,))
def test_initialization_eagle3_enforce_eager(self):
self.vllm_config.speculative_config.method = "eagle3"
self.vllm_config.speculative_config.draft_model_config.get_hidden_size.return_value = 2048
self.vllm_config.speculative_config.draft_model_config.get_inputs_embeds_size.return_value = 2048
self.vllm_config.compilation_config.mode = CompilationMode.NONE
self.vllm_config.compilation_config.pass_config = MagicMock()
self.vllm_config.compilation_config.pass_config.enable_sp = False
self.vllm_config.model_config.enforce_eager = True
init_ascend_config(self.vllm_config)
with set_current_vllm_config(self.vllm_config):
proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.assertEqual(proposer.hidden_size, 2048)
self.assertFalse(proposer.use_cuda_graph)
expected_max_num_tokens = proposer.max_num_tokens
self.assertEqual(proposer.hidden_states.shape, (expected_max_num_tokens, 2048))
def test_initialization_eagle3_full_graph_async(self):
self.vllm_config.speculative_config.method = "eagle3"
self.vllm_config.speculative_config.draft_model_config.get_hidden_size.return_value = 2048
self.vllm_config.speculative_config.draft_model_config.get_inputs_embeds_size.return_value = 2048
self.vllm_config.compilation_config.mode = CompilationMode.VLLM_COMPILE
self.vllm_config.model_config.enforce_eager = False
self.vllm_config.speculative_config.enforce_eager = False
self.vllm_config.scheduler_config.async_scheduling = True
init_ascend_config(self.vllm_config)
with set_current_vllm_config(self.vllm_config):
proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.assertEqual(proposer.hidden_size, 2048)
self.assertTrue(proposer.use_cuda_graph)
expected_max_num_tokens = proposer.max_num_tokens
self.assertEqual(proposer.hidden_states.shape, (expected_max_num_tokens, 2048))
def test_initialization_mtp_full_graph_async(self):
self.vllm_config.speculative_config.method = "mtp"
self.vllm_config.speculative_config.draft_model_config.get_hidden_size.return_value = 2048
self.vllm_config.speculative_config.draft_model_config.get_inputs_embeds_size.return_value = 2048
self.vllm_config.compilation_config.mode = CompilationMode.VLLM_COMPILE
self.vllm_config.model_config.enforce_eager = False
self.vllm_config.speculative_config.enforce_eager = False
self.vllm_config.scheduler_config.async_scheduling = True
init_ascend_config(self.vllm_config)
with set_current_vllm_config(self.vllm_config):
proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.assertEqual(proposer.hidden_size, 2048)
self.assertTrue(proposer.use_cuda_graph)
expected_max_num_tokens = proposer.max_num_tokens
self.assertEqual(proposer.hidden_states.shape, (expected_max_num_tokens, 2048))
def test_initialization_draft_model(self):
self.vllm_config.speculative_config.method = "draft_model"
self.vllm_config.speculative_config.parallel_drafting = False
# TODO(klyzhenko-vadim): remove when target_tp != draft_tp will be supported.
self.vllm_config.speculative_config.draft_parallel_config.tensor_parallel_size = 1
self.vllm_config.speculative_config.target_parallel_config.tensor_parallel_size = 1
init_ascend_config(self.vllm_config)
with set_current_vllm_config(self.vllm_config):
proposer = AscendDraftModelProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.assertTrue(isinstance(proposer, DraftModelProposer))
self.assertFalse(proposer.pass_hidden_states_to_model)
self.assertTrue(proposer.needs_extra_input_slots)
@unittest.skip("Skip due to the changes in #7153, fix me later")
class TestEagleProposerLoadModel(TestBase):
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.speculative_config.method = "eagle"
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.num_speculative_tokens = 2
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(2)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
# Mock get_ascend_config to return a properly configured mock
self.mock_get_ascend_config = patch("vllm_ascend.utils.get_ascend_config")
mock_config = self.mock_get_ascend_config.start()
mock_ascend_config = MagicMock()
mock_ascend_config.enable_flashcomm2_parallel_size = 0
mock_ascend_config.enable_context_parallel = False
mock_ascend_config.enable_flashcomm1 = False
mock_ascend_config.enable_matmul_allreduce = False
mock_ascend_config.weight_nz_mode = 1
mock_ascend_config.enable_mlapo = True
mock_ascend_config.enable_fused_mc2 = 0
mock_ascend_config.msmonitor_use_daemon = False
mock_ascend_config.enable_transpose_kv_cache_by_block = True
mock_config.return_value = mock_ascend_config
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Set the current vllm config
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.proposer.parallel_drafting = False
def tearDown(self):
self.mock_get_ascend_config.stop()
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
# Clear the current vllm config
set_current_vllm_config(None)
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_layers_from_vllm_config")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_model")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_pp_group")
def test_load_model_pp1(self, mock_pp_group, mock_get_model, mock_get_layers):
mock_pp_group.return_value.world_size = 1
mock_target_layer1 = MagicMock()
mock_target_layer2 = MagicMock()
mock_draft_layer1 = MagicMock()
mock_draft_layer3 = MagicMock()
mock_get_layers.side_effect = [
{"layer1": mock_target_layer1, "layer2": mock_target_layer2},
{},
{},
{"layer1": mock_draft_layer1, "layer3": mock_draft_layer3},
]
weight = torch.zeros(0)
mock_model = MagicMock()
mock_model.supports_multimodal = False
mock_model.lm_head = MagicMock()
mock_model.multimodal_cpu_fields = None
mock_model.merge_by_field_config = None
mock_model.model.embed_tokens = MagicMock()
mock_model.model.embed_tokens.weight = weight
mock_get_model.return_value = MagicMock()
mock_get_model.return_value.model.embed_tokens.weight = weight
with set_current_vllm_config(self.vllm_config):
self.proposer.load_model(mock_model)
mock_get_model.assert_called_once()
self.assertEqual(self.proposer.attn_layer_names, ["layer3"])
self.assertIs(self.proposer.model.model.embed_tokens, mock_model.model.embed_tokens)
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_layers_from_vllm_config")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_model")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_pp_group")
def test_load_model_pp_gt1(self, mock_pp_group, mock_get_model, mock_get_layers):
mock_pp_group.return_value.world_size = 2
mock_target_layer1 = MagicMock()
mock_draft_layer2 = MagicMock()
mock_get_layers.side_effect = [{"layer1": mock_target_layer1}, {}, {}, {"layer2": mock_draft_layer2}]
mock_model = MagicMock()
original_embed = MagicMock()
mock_model.multimodal_cpu_fields = None
mock_model.merge_by_field_config = None
mock_get_model.return_value = MagicMock(model=MagicMock(embed_tokens=original_embed))
with set_current_vllm_config(self.vllm_config):
self.proposer.load_model(mock_model)
self.assertIsNot(self.proposer.model.model.embed_tokens, mock_model.model.embed_tokens)
self.assertEqual(self.proposer.attn_layer_names, ["layer2"])
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_layers_from_vllm_config")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_model")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_pp_group")
@patch("vllm_ascend.spec_decode.llm_base_proposer.supports_multimodal")
def test_load_model_multimodal(self, mock_supports_multi, mock_pp_group, mock_get_model, mock_get_layers):
mock_model = MagicMock()
mock_model.get_language_model.return_value.lm_head = MagicMock()
mock_supports_multi.return_value = True
original_embed = MagicMock()
mock_get_model.return_value = MagicMock(model=MagicMock(embed_tokens=original_embed))
mock_target_layer1 = MagicMock()
mock_draft_layer2 = MagicMock()
mock_get_layers.side_effect = [{"layer1": mock_target_layer1}, {}, {}, {"layer2": mock_draft_layer2}]
mock_pp_group.return_value.world_size = 2
self.proposer.model = MagicMock()
with set_current_vllm_config(self.vllm_config):
self.proposer.load_model(mock_model)
self.assertEqual(mock_model.get_language_model.call_count, 2)
self.assertIs(self.proposer.model.lm_head, mock_model.get_language_model.return_value.lm_head)
class TestEagleProposerDummyRun(TestBase):
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.speculative_config.num_speculative_tokens = 4
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.runner.pin_memory = False
self.runner._sync_metadata_across_dp.return_value = (8, torch.tensor([8]), CUDAGraphMode.NONE)
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.model_config.use_mla = False
self.vllm_config.model_config.hf_text_config = MagicMock(
spec=[]
) # Empty spec to prevent hasattr from returning True
self.vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.parallel_config.pipeline_parallel_size = 1
self.vllm_config.model_config.enforce_eager = True
self.vllm_config.model_config.is_deepseek_mla = False
self.vllm_config.kv_transfer_config = None
self.vllm_config.compilation_config = MagicMock()
self.vllm_config.compilation_config.pass_config = MagicMock()
self.vllm_config.compilation_config.pass_config.enable_sp = False
self.vllm_config.cache_config = MagicMock()
self.vllm_config.cache_config.block_size = 16
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(4)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
# Mock get_ascend_config to return a properly configured mock
self.mock_get_ascend_config = patch("vllm_ascend.utils.get_ascend_config")
mock_config = self.mock_get_ascend_config.start()
mock_ascend_config = MagicMock()
mock_ascend_config.enable_flashcomm2_parallel_size = 0
mock_ascend_config.enable_context_parallel = False
mock_ascend_config.enable_flashcomm1 = False
mock_ascend_config.enable_matmul_allreduce = False
mock_ascend_config.weight_nz_mode = 1
mock_ascend_config.enable_mlapo = True
mock_ascend_config.enable_fused_mc2 = 0
mock_ascend_config.msmonitor_use_daemon = False
mock_ascend_config.enable_transpose_kv_cache_by_block = True
mock_config.return_value = mock_ascend_config
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Mock parallel state functions
self.mock_tp_world_size = patch(
"vllm_ascend.ascend_forward_context.get_tensor_model_parallel_world_size", return_value=1
)
self.mock_tp_world_size.start()
mock_dp_group = MagicMock()
mock_dp_group.world_size = 1
self.mock_dp_group = patch("vllm_ascend.ascend_forward_context.get_dp_group", return_value=mock_dp_group)
self.mock_dp_group.start()
# Set the current vllm config
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.proposer.model = MagicMock()
self.proposer._runnable = MagicMock()
self.proposer.update_stream = MagicMock()
def tearDown(self):
self.mock_get_ascend_config.stop()
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
self.mock_tp_world_size.stop()
self.mock_dp_group.stop()
# Clear the current vllm config
set_current_vllm_config(None)
# cpu does not support parallel-group, let alone `sp`
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch(
"vllm_ascend.spec_decode.llm_base_proposer.get_forward_context", **{"return_value.flash_comm_v1_enabled": False}
)
@patch("vllm_ascend.spec_decode.llm_base_proposer.set_ascend_forward_context")
def test_dummy_run_basic(self, mock_context, mock_get_context, mock_get_context_2):
num_tokens = 32
with_prefill = False
# cpu does not support `torch.ops.vllm.maybe_pad_and_reduce`
with set_current_vllm_config(self.vllm_config):
self.proposer.enable_shared_expert_dp = False
self.proposer.dummy_run(num_tokens=num_tokens, with_prefill=with_prefill)
self.assertTrue(self.proposer._runnable.call_count == 1)
# cpu does not support parallel-group, let alone `sp`
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch(
"vllm_ascend.spec_decode.llm_base_proposer.get_forward_context", **{"return_value.flash_comm_v1_enabled": False}
)
@patch("vllm_ascend.spec_decode.llm_base_proposer.set_ascend_forward_context")
def test_dummy_run_with_prefill(self, mock_context, mock_get_context, mock_get_context_2):
mock_context.return_value.__enter__.return_value = None
# cpu does not support `torch.ops.vllm.maybe_pad_and_reduce`
with set_current_vllm_config(self.vllm_config):
self.proposer.enable_shared_expert_dp = False
self.proposer.dummy_run(num_tokens=64, with_prefill=True, num_reqs=4)
self.assertTrue(self.proposer._runnable.call_count == 1)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.spec_decode.llm_base_proposer.update_full_graph_params")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_forward_context")
@patch("vllm_ascend.spec_decode.llm_base_proposer.set_ascend_forward_context")
def test_dummy_run_in_graph_capture(
self, mock_context, mock_get_context, mock_update_full_graph_params, mock_get_context_2
):
last_use_cuda_graph = self.proposer.use_cuda_graph
mock_return_context = MagicMock()
mock_return_context.cudagraph_runtime_mode = CUDAGraphMode.FULL
mock_return_context.capturing = True
# cpu does not support parallel-group, let alone `sp`
mock_return_context.flash_comm_v1_enabled = False
mock_get_context.return_value = mock_return_context
mock_get_context_2.return_value = mock_return_context
self.proposer.use_cuda_graph = True
# cpu does not support `torch.ops.vllm.maybe_pad_and_reduce`
with set_current_vllm_config(self.vllm_config):
self.proposer.enable_shared_expert_dp = False
self.proposer.dummy_run(num_tokens=64, in_graph_capturing=True, aclgraph_runtime_mode=CUDAGraphMode.FULL)
self.assertTrue(self.proposer._runnable.call_count == 1)
mock_update_full_graph_params.assert_not_called()
self.proposer.use_cuda_graph = last_use_cuda_graph
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.spec_decode.llm_base_proposer.update_full_graph_params")
@patch("vllm_ascend.spec_decode.llm_base_proposer.get_forward_context")
@patch("vllm_ascend.spec_decode.llm_base_proposer.set_ascend_forward_context")
def test_dummy_run_in_graph_run(
self, mock_context, mock_get_context, mock_update_full_graph_params, mock_get_context_2
):
last_use_cuda_graph = self.proposer.use_cuda_graph
mock_return_context = MagicMock()
mock_return_context.cudagraph_runtime_mode = CUDAGraphMode.FULL
mock_return_context.capturing = False
# cpu does not support parallel-group, let alone `sp`
mock_return_context.flash_comm_v1_enabled = False
mock_get_context.return_value = mock_return_context
mock_get_context_2.return_value = mock_return_context
self.proposer.use_cuda_graph = True
self.proposer.draft_attn_groups = [MagicMock()]
# cpu does not support `torch.ops.vllm.maybe_pad_and_reduce`
with set_current_vllm_config(self.vllm_config):
self.proposer.enable_shared_expert_dp = False
self.proposer.dummy_run(num_tokens=64, in_graph_capturing=False, aclgraph_runtime_mode=CUDAGraphMode.FULL)
self.assertTrue(self.proposer._runnable.call_count == 1)
self.assertTrue(mock_update_full_graph_params.call_count == 1)
self.proposer.use_cuda_graph = last_use_cuda_graph
class TestEagleProposerHelperMethods(TestBase):
# TODO: Can add some tests about prepare_next_token_ids in future.
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.scheduler_config = MagicMock(max_num_seqs=3)
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.input_batch = MagicMock()
self.runner.input_batch.req_ids = [0, 1, 2]
self.runner.arange_np = np.arange(10)
self.runner.input_batch.num_reqs = 3
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.num_speculative_tokens = 2
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(2)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Set the current vllm config
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
def tearDown(self):
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
# Clear the current vllm config
set_current_vllm_config(None)
# TODO: This is equivalent to disable_padded_drafter_batch=True.
# We need to add a test_prepare_inputs_padded in future.
def test_prepare_inputs(self):
self.proposer.token_arange_np = np.arange(10)
mock_attn = MagicMock()
mock_attn.slot_mapping = torch.tensor([0, 1, 2, 3, 4, 5])
num_rejected = torch.tensor([1, 0, 1], device=self.device)
mock_return_attn = MagicMock()
with (
set_current_vllm_config(self.vllm_config),
patch.object(self.proposer, "prepare_inputs", return_value=(mock_return_attn, torch.tensor([1, 2, 4]))),
):
return_attn, indices = self.proposer.prepare_inputs(mock_attn, num_rejected)
self.assertEqual(indices.tolist(), [1, 2, 4])
# fmt: off
class TestEagleProposerMaybePadAndGather:
@pytest.fixture(autouse=True)
def setUp_and_tearDown(self):
self.check_mock()
self.device = torch.device("npu")
yield
def _new_proposer(
self,
method,
*,
is_multimodal_model=False,
enable_shared_expert_dp=False,
):
proposer = object.__new__(AscendEagleProposer)
proposer.method = method
proposer.is_multimodal_model = is_multimodal_model
proposer.enable_shared_expert_dp = enable_shared_expert_dp
return proposer
def _extra_ctx(self, flash_comm_v1_enabled):
extra_ctx = MagicMock()
extra_ctx.flash_comm_v1_enabled = flash_comm_v1_enabled
return extra_ctx
@pytest.mark.parametrize(
"flash_comm_v1_enabled,is_multimodal_model,expect_reduce",
[
(True, False, True),
(False, False, False),
(True, True, False),
],
)
def test_mtp_maybe_pad_and_reduce(
self,
flash_comm_v1_enabled,
is_multimodal_model,
expect_reduce,
):
proposer = self._new_proposer("mtp", is_multimodal_model=is_multimodal_model)
model_hidden_states = torch.arange(12, device=self.device, dtype=torch.float32).view(6, 2)
model_positions = torch.arange(6, device=self.device, dtype=torch.int64)
def fake_pad_and_reduce(input_tensor):
return input_tensor[::2].contiguous()
with (
patch("vllm_ascend.spec_decode.llm_base_proposer._EXTRA_CTX", new=self._extra_ctx(flash_comm_v1_enabled)),
patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=fake_pad_and_reduce, create=True) as mock_reduce,
):
reduced_hidden_states, reduced_positions = proposer.maybe_pad_and_reduce(
model_hidden_states, model_positions
)
if expect_reduce:
assert mock_reduce.call_count == 2
assert mock_reduce.call_args_list[0].args[0].shape == (6, 2)
assert mock_reduce.call_args_list[1].args[0].shape == (6, 1)
assert torch.equal(reduced_hidden_states, model_hidden_states[::2])
assert torch.equal(reduced_positions, model_positions[::2])
else:
mock_reduce.assert_not_called()
assert reduced_hidden_states is model_hidden_states
assert reduced_positions is model_positions
@pytest.mark.parametrize(
"flash_comm_v1_enabled,expect_split",
[
(True, True),
(False, False),
],
)
def test_eagle_maybe_pad_and_reduce(
self,
flash_comm_v1_enabled,
expect_split,
):
proposer = self._new_proposer("eagle3")
model_hidden_states = torch.arange(12, device=self.device, dtype=torch.float32).view(6, 2)
model_positions = torch.arange(6, device=self.device, dtype=torch.int64)
tp_group = MagicMock()
tp_group.world_size = 2
tp_group.rank = 1
with (
patch("vllm_ascend.spec_decode.llm_base_proposer._EXTRA_CTX", new=self._extra_ctx(flash_comm_v1_enabled)),
patch("vllm_ascend.spec_decode.llm_base_proposer.get_tp_group", return_value=tp_group) as mock_get_tp_group,
):
reduced_hidden_states, reduced_positions = proposer.maybe_pad_and_reduce(
model_hidden_states, model_positions
)
if expect_split:
expected_hidden_states = torch.tensor([[6.0, 7.0], [8.0, 9.0], [10.0, 11.0]], device=self.device)
mock_get_tp_group.assert_called_once()
assert reduced_hidden_states.shape == (3, 2)
assert torch.equal(reduced_hidden_states, expected_hidden_states)
assert torch.equal(model_hidden_states[:3], expected_hidden_states)
else:
mock_get_tp_group.assert_not_called()
assert reduced_hidden_states is model_hidden_states
assert reduced_positions is model_positions
@pytest.mark.parametrize(
"enable_shared_expert_dp,hidden_states_is_none,expect_gather",
[
(True, False, True),
(True, True, True),
(False, False, False),
],
)
def test_mtp_maybe_all_gather_and_unpad(
self,
enable_shared_expert_dp,
hidden_states_is_none,
expect_gather,
):
proposer = self._new_proposer("mtp", enable_shared_expert_dp=enable_shared_expert_dp)
last_hidden_states = torch.arange(6, device=self.device, dtype=torch.float32).view(3, 2)
positions = torch.tensor([10, 11, 12], device=self.device, dtype=torch.int64)
hidden_states = None if hidden_states_is_none else last_hidden_states + 1000
def fake_all_gather_and_unpad(input_tensor, label):
assert label is True
return torch.cat((input_tensor, input_tensor + 100), dim=0)
with patch(
"torch.ops.vllm.maybe_all_gather_and_maybe_unpad",
side_effect=fake_all_gather_and_unpad,
create=True,
) as mock_all_gather:
gathered_last_hidden_states, gathered_positions, gathered_hidden_states = (
proposer.maybe_all_gather_and_unpad(last_hidden_states, positions, hidden_states)
)
if expect_gather:
expected_last_hidden_states = torch.cat((last_hidden_states, last_hidden_states + 100), dim=0)
expected_positions = torch.cat((positions, positions + 100), dim=0)
assert mock_all_gather.call_count == 2
assert torch.equal(gathered_last_hidden_states, expected_last_hidden_states)
assert torch.equal(gathered_positions, expected_positions)
if hidden_states_is_none:
assert gathered_hidden_states is None
else:
assert gathered_hidden_states is gathered_last_hidden_states
else:
mock_all_gather.assert_not_called()
assert gathered_last_hidden_states is last_hidden_states
assert gathered_positions is positions
assert gathered_hidden_states is hidden_states
@pytest.mark.parametrize(
"flash_comm_v1_enabled,hidden_states_is_none,expect_gather",
[
(True, False, True),
(True, True, True),
(False, False, False),
],
)
def test_eagle_maybe_all_gather_and_unpad(
self,
flash_comm_v1_enabled,
hidden_states_is_none,
expect_gather,
):
proposer = self._new_proposer("eagle3")
last_hidden_states = torch.arange(6, device=self.device, dtype=torch.float32).view(3, 2)
positions = torch.tensor([10, 11, 12], device=self.device, dtype=torch.int64)
hidden_states = None if hidden_states_is_none else last_hidden_states + 1000
def fake_all_gather_and_unpad(input_tensor, label):
assert label is True
return torch.cat((input_tensor, input_tensor + 100), dim=0)
with (
patch("vllm_ascend.spec_decode.llm_base_proposer._EXTRA_CTX", new=self._extra_ctx(flash_comm_v1_enabled)),
patch(
"torch.ops.vllm.maybe_all_gather_and_maybe_unpad",
side_effect=fake_all_gather_and_unpad,
create=True,
) as mock_all_gather,
):
gathered_last_hidden_states, gathered_positions, gathered_hidden_states = (
proposer.maybe_all_gather_and_unpad(last_hidden_states, positions, hidden_states)
)
if expect_gather:
expected_last_hidden_states = torch.cat((last_hidden_states, last_hidden_states + 100), dim=0)
assert torch.equal(gathered_last_hidden_states, expected_last_hidden_states)
assert gathered_positions is positions
if hidden_states_is_none:
assert mock_all_gather.call_count == 1
assert gathered_hidden_states is None
else:
expected_hidden_states = torch.cat((hidden_states, hidden_states + 100), dim=0) # type: ignore
assert mock_all_gather.call_count == 2
assert torch.equal(gathered_hidden_states, expected_hidden_states)
else:
mock_all_gather.assert_not_called()
assert gathered_last_hidden_states is last_hidden_states
assert gathered_positions is positions
assert gathered_hidden_states is hidden_states
def check_mock(self):
import vllm_ascend.spec_decode.llm_base_proposer
assert hasattr(vllm_ascend.spec_decode.llm_base_proposer, "AscendSpecDecodeBaseProposer")
RunnerCls = vllm_ascend.spec_decode.llm_base_proposer.AscendSpecDecodeBaseProposer
assert hasattr(RunnerCls, "maybe_pad_and_reduce")
sig = inspect.signature(RunnerCls.maybe_pad_and_reduce)
assert self.get_param_names(sig) == ["self", "hidden_states", "positions"]
assert hasattr(RunnerCls, "maybe_all_gather_and_unpad")
sig = inspect.signature(RunnerCls.maybe_all_gather_and_unpad)
assert self.get_param_names(sig) == ["self", "last_hidden_states", "positions", "hidden_states"]
def get_param_names(self, sig):
return [p.name for p in sig.parameters.values()]
# fmt: on
# fmt: off
class TestEagleProposerPropose:
@pytest.fixture(autouse=True)
def setUp_and_tearDown(self):
# before mock and patch, add assertions to ensure
# that the mocked functions and parameters exist
self.check_mock()
clear_ascend_config()
self.mock_get_ascend_config = patch("vllm_ascend.utils.get_ascend_config")
mock_get_ascend_config = self.mock_get_ascend_config.start()
mock_ascend_config = MagicMock()
mock_ascend_config.enable_flashcomm2_parallel_size = 0
mock_ascend_config.enable_context_parallel = False
mock_ascend_config.enable_flashcomm1 = False
mock_ascend_config.enable_matmul_allreduce = False
mock_ascend_config.weight_nz_mode = 1
mock_ascend_config.enable_mlapo = True
mock_ascend_config.enable_fused_mc2 = 0
mock_ascend_config.msmonitor_use_daemon = False
mock_ascend_config.enable_transpose_kv_cache_by_block = True
mock_get_ascend_config.return_value = mock_ascend_config
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.use_v2_model_runner = False
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.speculative_config.num_speculative_tokens = 3
self.vllm_config.speculative_config.method = "eagle3"
self.vllm_config.speculative_config.parallel_drafting = False
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.max_num_tokens = 8192
self.runner.max_num_reqs = 256
self.runner.pin_memory = False
self.runner.pcp_manager = None
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 32768
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.model_config.use_mla = False
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.parallel_config.pipeline_parallel_size = 1
self.vllm_config.model_config.enforce_eager = True
self.vllm_config.model_config.is_deepseek_mla = False
self.vllm_config.kv_transfer_config = None
self.vllm_config.compilation_config = MagicMock()
self.vllm_config.compilation_config.pass_config = MagicMock()
self.vllm_config.compilation_config.pass_config.enable_sp = False
self.vllm_config.cache_config = MagicMock()
self.vllm_config.cache_config.block_size = 16
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(4)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
self.ascend_config = init_ascend_config(self.vllm_config)
self.ascend_config.enable_flashcomm2_parallel_size = 0
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Mock parallel state functions
self.mock_tp_world_size = patch(
"vllm_ascend.ascend_forward_context.get_tensor_model_parallel_world_size", return_value=1
)
self.mock_tp_world_size.start()
mock_dp_group = MagicMock()
mock_dp_group.world_size = 1
self.mock_dp_group = patch("vllm_ascend.ascend_forward_context.get_dp_group", return_value=mock_dp_group)
self.mock_dp_group.start()
# Mock sp
self.mock_enable_sp = patch(
"vllm_ascend.utils.enable_sp", return_value=False
)
self.mock_enable_sp.start()
# Set the current vllm config
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
yield
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
self.mock_tp_world_size.stop()
self.mock_dp_group.stop()
self.mock_get_ascend_config.stop()
# Clear the current vllm config
set_current_vllm_config(None)
clear_ascend_config()
# config: prefill and decode, Qwen3-8B, tp1, enforce_eager, no_async_scheduling, eagle3, k=3, "disable_padded_drafter_batch": False
@pytest.mark.parametrize(
'flag_prefill_decode, query_start_loc, query_start_loc_cpu, seq_lens, num_reqs,' \
'num_actual_tokens, max_query_len, max_seq_len, block_table_tensor,' \
'slot_mapping, causal, logits_indices_padded, num_logits_indices,' \
'encoder_seq_lens, encoder_seq_lens_cpu, dcp_local_seq_lens,' \
'dcp_local_seq_lens_cpu, _seq_lens_cpu, _num_computed_tokens_cpu,' \
'_num_computed_tokens_cache, seq_lens_cpu, num_computed_tokens_cpu,' \
'decode_token_per_req, actual_seq_lengths_q, positions, attn_state,' \
'graph_pad_size, num_input_tokens, prefill_context_parallel_metadata',
[
(
"prefill", torch.tensor([ 0, 13], device=torch.device("cpu"), dtype=torch.int32), torch.tensor([ 0, 13], dtype=torch.int32),
torch.tensor([13], device=torch.device("cpu"), dtype=torch.int32), 1, 13, 13, 13,
torch.eye(256, device=torch.device("cpu"), dtype=torch.int32)[0].unsqueeze(0),
torch.tensor([128, 129, 130, 131, 132, 133, 134, 135, 136, 137, 138, 139, 140], device=torch.device("cpu"), dtype=torch.int32),
True, None, None, None, None, None, None, None, None, None, torch.tensor([13], dtype=torch.int32), torch.tensor([0], dtype=torch.int32), 4, [],
torch.cat([torch.arange(13), torch.zeros(8704 - 13)]),
AscendAttentionState.PrefillNoCache, -1, 13, None
),
(
"decode", torch.tensor([ 0, 4, 8, 12], device=torch.device("cpu"), dtype=torch.int32), torch.tensor([ 0, 4, 8, 12], dtype=torch.int32),
torch.tensor([21, 17, 17], device=torch.device("cpu"), dtype=torch.int32), 3, 12, 4, 0,
torch.cat([torch.eye(256, device="cpu", dtype=torch.int32)[0].unsqueeze(0)*i for i in [1,2,3]], dim=0),
torch.tensor([145, 146, 147, 148, 269, 270, 271, 272, 397, 398, 399, 400], device=torch.device("cpu"), dtype=torch.int32),
True, None, None, None, None, None, None, None, None, None, torch.tensor([21, 17, 17], dtype=torch.int32), torch.tensor([17, 13, 13], dtype=torch.int32), 4, [],
torch.cat([torch.tensor([17, 18, 19, 20, 13, 14, 15, 16, 13, 14, 15, 16, 8, 9, 10, 11, 12, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]), torch.zeros(8704 - 30)]),
AscendAttentionState.ChunkedPrefill, -1, 12, None
),
(
"decode_and_prefill", torch.tensor([ 0, 4, 17, 30], device=torch.device("cpu"), dtype=torch.int32), torch.tensor([ 0, 4, 17, 30], dtype=torch.int32),
torch.tensor([17, 13, 13], device=torch.device("cpu"), dtype=torch.int32), 3, 30, 13, 0,
torch.cat([torch.eye(256, device="cpu", dtype=torch.int32)[0].unsqueeze(0)*i for i in [1,2,3]], dim=0),
torch.tensor([141, 142, 143, 144, 256, 257, 258, 259, 260, 261, 262, 263, 264, 265, 266, 267, 268, 384, 385, 386, 387, 388, 389, 390, 391, 392, 393, 394,
395, 396], device=torch.device("cpu"), dtype=torch.int32),
True, None, None, None, None, None, None, torch.tensor([17, 13, 13], dtype=torch.int32), None, None, torch.tensor([17, 13, 13], dtype=torch.int32),
torch.tensor([13, 0, 0], dtype=torch.int32), 4, [],
torch.cat([torch.tensor([13, 14, 15, 16, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]), torch.zeros(8704 - 30)]),
AscendAttentionState.ChunkedPrefill, -1, 30, None
),
]
)
# config: prefill and decode, Qwen3-30B, tp2, ep_enable, enforce_eager, no_async_scheduling, eagle3, k=3, "disable_padded_drafter_batch": False
@pytest.mark.parametrize('model_type', ['qwen_dense','qwen_moe', 'deepseek'])
@pytest.mark.parametrize('graphmode', ['eager','full'])
@patch('vllm_ascend.spec_decode.eagle_proposer.AscendEagleProposer.get_model')
def test_propose(self, mock_get_model, graphmode, model_type, flag_prefill_decode,
query_start_loc, query_start_loc_cpu, seq_lens, num_reqs,
num_actual_tokens, max_query_len, max_seq_len, block_table_tensor,
slot_mapping, causal, logits_indices_padded, num_logits_indices,
encoder_seq_lens, encoder_seq_lens_cpu, dcp_local_seq_lens,
dcp_local_seq_lens_cpu, _seq_lens_cpu, _num_computed_tokens_cpu,
_num_computed_tokens_cache, seq_lens_cpu, num_computed_tokens_cpu,
decode_token_per_req, actual_seq_lengths_q, positions, attn_state,
graph_pad_size, num_input_tokens, prefill_context_parallel_metadata
):
# adjust for fullgraph mode
if graphmode == 'full':
if model_type == "qwen_dense" and self.is_decode(flag_prefill_decode):
slot_mapping = torch.tensor([145, 146, 147, 148, 269, 270, 271, 272, 397, 398, 399, 400, -1, -1], device=torch.device("cpu"), dtype=torch.int32)
positions = torch.cat([torch.tensor([17, 18, 19, 20, 13, 14, 15, 16, 13, 14, 15, 16, 0, 0, 0, 0, 12, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]), torch.zeros(8704 - 30)])
num_input_tokens = 16
target_model_batch_desc = BatchDescriptor(num_tokens=num_input_tokens, num_reqs=4, uniform=True, has_lora=False, num_active_loras=0)
self.proposer.use_cuda_graph = True
else:
pytest.skip("For the entire graph test, only one model needs to be tested to avoid repeated tests.")
if flag_prefill_decode == 'decode_and_prefill':
if not (model_type == 'qwen_dense' and graphmode == 'eager'):
pytest.skip(
"decode_and_prefill case only need test once"
)
# mock and adjust functions and var in propose
if model_type == 'deepseek':
self.proposer.method = 'mtp'
if not self.is_decode(flag_prefill_decode):
num_actual_tokens = 9
self.runner._sync_metadata_across_dp.return_value = (num_actual_tokens, None, False)
self.proposer.model = MagicMock(spec=Eagle3LlamaForCausalLM)
custom_combined_hidden_states = torch.zeros(num_actual_tokens, 4096, device=self.device, dtype=torch.bfloat16)
self.proposer.model.combine_hidden_states.return_value = custom_combined_hidden_states
mock_get_model.return_value = self.proposer.model
self.proposer.hidden_size = 4096
if model_type == 'deepseek':
self.proposer.hidden_states = torch.zeros(8192, 7168, device=self.device, dtype=torch.bfloat16)
else:
self.proposer.hidden_states = torch.zeros(8192, 4096, device=self.device, dtype=torch.bfloat16)
mock_attn_group = MagicMock()
mock_builder = MagicMock()
mock_attn_metadata = MagicMock()
mock_builder.build.return_value = mock_attn_metadata
mock_attn_group.get_metadata_builder.return_value = mock_builder
self.proposer.draft_attn_groups = [mock_attn_group]
self.proposer.attn_layer_names = ['model.layers.36.self_attn.attn']
self.proposer.kernel_block_size = 128
self.proposer.block_size = 128
self.proposer._runnable = MagicMock()
self.proposer._runnable.return_value = [0, 0, 0]
captured_common_attn_metadata = None
original_method = self.proposer.attn_update_stack_num_spec_norm
mock_bd = MagicMock()
mock_bd.num_tokens = 16
self.proposer.query_start_loc = MagicMock()
self.proposer.query_start_loc.gpu = torch.tensor([0, 4, 8, 12, 16], device=torch.device("cpu"), dtype=torch.int32)
self.proposer.query_start_loc.cpu = torch.tensor([0, 4, 8, 12, 16], device=torch.device("cpu"), dtype=torch.int32)
self.runner.cudagraph_dispatcher.dispatch.return_value = (CUDAGraphMode.FULL, mock_bd)
self.runner._pad_query_start_loc_for_fia.return_value = 4
self.runner.query_start_loc.gpu = torch.tensor([0, 4, 8, 12, 16], device=torch.device("cpu"), dtype=torch.int32)
self.runner.query_start_loc.cpu = torch.tensor([0, 4, 8, 12, 16], device=torch.device("cpu"), dtype=torch.int32)
self.runner.seq_lens = seq_lens
self.runner.optimistic_seq_lens_cpu = seq_lens_cpu
self.proposer._update_full_graph_params = MagicMock()
def side_effect(*args, **kwargs):
nonlocal captured_common_attn_metadata
res_common, res_attn = original_method(*args, **kwargs)
captured_common_attn_metadata = res_common
return res_common, res_attn
# create common_attn_metadata
mock_common_attn_metadata= MagicMock()
if not self.is_decode(flag_prefill_decode):
if flag_prefill_decode == 'prefill':
mock_common_attn_metadata.batch_size.return_value = 1
if flag_prefill_decode == 'decode_and_prefill':
mock_common_attn_metadata.batch_size.return_value = 3
if model_type == 'qwen_moe':
_seq_lens_cpu = torch.tensor([13], dtype=torch.int32)
if model_type == 'deepseek':
query_start_loc = torch.tensor([0, 9], device=torch.device("cpu"), dtype=torch.int32)
query_start_loc_cpu = torch.tensor([0, 9], device=torch.device("cpu"), dtype=torch.int32)
seq_lens = torch.tensor([9], device=torch.device("cpu"), dtype=torch.int32)
max_query_len = 9
max_seq_len = 9
slot_mapping = torch.tensor([128, 129, 130, 131, 132, 133, 134, 135, 136], device=torch.device("cpu"), dtype=torch.int32)
_seq_lens_cpu = torch.tensor([9], dtype=torch.int32)
seq_lens_cpu = torch.tensor([9], dtype=torch.int32)
positions = torch.cat([torch.arange(9), torch.zeros(8704 - 9)])
num_input_tokens = 9
if self.is_decode(flag_prefill_decode):
mock_common_attn_metadata.batch_size.return_value = 3
if model_type == 'qwen_moe':
seq_lens = torch.tensor([19, 17, 17], device=torch.device("cpu"), dtype=torch.int32)
slot_mapping = torch.tensor([143, 144, 145, 146, 269, 270, 271, 272, 397, 398, 399, 400], device=torch.device("cpu"), dtype=torch.int32)
seq_lens_cpu = torch.tensor([19, 17, 17], dtype=torch.int32)
num_computed_tokens_cpu = torch.tensor([15, 13, 13], dtype=torch.int32)
positions = torch.cat([torch.tensor([15, 16, 17, 18, 13, 14, 15, 16, 13, 14, 15, 16, 8, 9, 10, 11, 12, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]), torch.zeros(8704 - 30)])
if model_type == 'deepseek':
seq_lens = torch.tensor([14, 13, 14], device=torch.device("cpu"), dtype=torch.int32)
slot_mapping = torch.tensor([138, 139, 140, 141, 265, 266, 267, 268, 394, 395, 396, 397], device=torch.device("cpu"), dtype=torch.int32)
seq_lens_cpu = torch.tensor([14, 13, 14], dtype=torch.int32)
num_computed_tokens_cpu = torch.tensor([10, 9, 10], dtype=torch.int32)
positions = torch.cat([torch.tensor([10, 11, 12, 13, 9, 10, 11, 12, 10, 11, 12, 13, 8, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9]), torch.zeros(8704 - 23)])
attn_state = AscendAttentionState.SpecDecoding
self.value_mock_common_attn_metadata(mock_common_attn_metadata, query_start_loc, query_start_loc_cpu, seq_lens, num_reqs,
num_actual_tokens, max_query_len, max_seq_len, block_table_tensor,
slot_mapping, causal, logits_indices_padded, num_logits_indices,
encoder_seq_lens, encoder_seq_lens_cpu, dcp_local_seq_lens,
dcp_local_seq_lens_cpu, _seq_lens_cpu, _num_computed_tokens_cpu,
_num_computed_tokens_cache, seq_lens_cpu, num_computed_tokens_cpu,
decode_token_per_req, actual_seq_lengths_q, positions, attn_state,
graph_pad_size, num_input_tokens, prefill_context_parallel_metadata
)
# create other parameters
if not self.is_decode(flag_prefill_decode):
if model_type == 'qwen_dense' or model_type == 'qwen_moe':
if flag_prefill_decode == 'prefill':
target_token_ids = torch.tensor([151644, 872, 198, 5501, 7512, 14678, 51765, 30, 151645, 198, 151644, 77091, 198], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], device=self.device)
next_token_ids = torch.tensor([151667], device=self.device, dtype=torch.int32)
req_scheduled_tokens = {'0-8222703c': 13}
if flag_prefill_decode == 'decode_and_prefill':
target_token_ids = torch.tensor([151667, 198, 32313, 11, 151644, 872, 198, 5501, 7512, 387, 23649, 30, 151645, 198, 151644, 77091, 198, 151644,
872, 198, 5501, 7512, 557, 30070, 30, 151645, 198, 151644, 77091, 198], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([13, 14, 15, 16, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], device=self.device)
next_token_ids = torch.tensor([279, 151667, 151667], device=self.device, dtype=torch.int32)
req_scheduled_tokens = {'0-b0f8a3bc': 4, '1-99a1b5fa': 13, '2-8a6a85d3': 13}
if model_type == 'deepseek':
target_token_ids = torch.tensor([ 0, 0, 128803, 12473, 9734, 19991, 50096, 33, 128804], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8], device=self.device)
target_hidden_states = torch.zeros(num_actual_tokens, 7168, device=self.device, dtype=torch.bfloat16)
next_token_ids = torch.tensor([128798], device=self.device, dtype=torch.int32)
req_scheduled_tokens = {'0-b4ed8210': 9}
if model_type == 'qwen_dense':
if flag_prefill_decode == 'prefill':
target_hidden_states = torch.zeros(num_actual_tokens, 12288, device=self.device, dtype=torch.bfloat16)
if flag_prefill_decode == 'decode_and_prefill':
target_hidden_states = torch.zeros(30, 12288, device=self.device, dtype=torch.bfloat16)
if model_type == 'qwen_moe':
target_hidden_states = torch.zeros(num_actual_tokens, 6144, device=self.device, dtype=torch.bfloat16)
if flag_prefill_decode == 'decode_and_prefill':
token_indices_to_sample = torch.tensor([3, 16, 29], device=self.device, dtype=torch.int32)
else:
token_indices_to_sample = None
target_model_batch_desc = BatchDescriptor(num_tokens=num_actual_tokens, num_reqs=None, uniform=False, has_lora=False, num_active_loras=0)
mock_sampling_metadata = MagicMock()
mm_embed_inputs = None
long_seq_metadata = None
num_prefill_reqs = 0
num_decode_reqs = 0
scheduler_output = MagicMock()
num_scheduled_tokens = num_actual_tokens
if flag_prefill_decode == 'decode_and_prefill':
num_rejected_tokens_gpu = torch.tensor([0, 0, 0], device=self.device, dtype=torch.int32)
else:
num_rejected_tokens_gpu = None
if self.is_decode(flag_prefill_decode):
if model_type == 'qwen_dense':
target_token_ids = torch.tensor([279, 1196, 374, 8014, 151667, 198, 32313, 11, 151667, 198, 32313, 11], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([17, 18, 19, 20, 13, 14, 15, 16, 13, 14, 15, 16], device=self.device)
target_hidden_states = torch.zeros(num_actual_tokens, 12288, device=self.device, dtype=torch.bfloat16)
next_token_ids = torch.tensor([4588, 279, 279], device=self.device, dtype=torch.int32)
token_indices_to_sample = torch.tensor([1, 7, 11], device=self.device, dtype=torch.int32)
num_rejected_tokens_gpu = torch.tensor([2, 0, 0], device=self.device, dtype=torch.int32)
if model_type == 'qwen_moe':
target_token_ids = torch.tensor([32313, 2776, 198, 198, 151667, 198, 198, 198, 151667, 198, 198, 198], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([15, 16, 17, 18, 13, 14, 15, 16, 13, 14, 15, 16], device=self.device)
target_hidden_states = torch.zeros(num_actual_tokens, 6144, device=self.device, dtype=torch.bfloat16)
next_token_ids = torch.tensor([11, 32313, 32313], device=self.device, dtype=torch.int32)
token_indices_to_sample = torch.tensor([0, 5, 9], device=self.device, dtype=torch.int32)
num_rejected_tokens_gpu = torch.tensor([3, 2, 2], device=self.device, dtype=torch.int32)
if model_type == 'deepseek':
target_token_ids = torch.tensor([201, 33001, 14, 832, 128798, 271, 5, 128798, 128798, 271, 5, 128798], device=self.device, dtype=torch.int32)
target_positions = torch.tensor([10, 11, 12, 13, 9, 10, 11, 12, 10, 11, 12, 13], device=self.device)
target_hidden_states = torch.zeros(num_actual_tokens, 7168, device=self.device, dtype=torch.bfloat16)
next_token_ids = torch.tensor([270, 128799, 201], device=self.device, dtype=torch.int32)
token_indices_to_sample = torch.tensor([2, 5, 8], device=self.device, dtype=torch.int32)
num_rejected_tokens_gpu = torch.tensor([1, 2, 3], device=self.device, dtype=torch.int32)
target_model_batch_desc = BatchDescriptor(num_tokens=num_actual_tokens, num_reqs=None, uniform=False, has_lora=False, num_active_loras=0)
mock_sampling_metadata = MagicMock()
mm_embed_inputs = None
req_scheduled_tokens = {'0-b69afbe5': 4, '1-b60368b9': 4, '2-82281e95': 4}
long_seq_metadata = None
num_prefill_reqs = 0
num_decode_reqs = 0
scheduler_output = MagicMock()
num_scheduled_tokens = num_actual_tokens
#run
with (
patch.object(self.proposer, 'attn_update_stack_num_spec_norm', side_effect=side_effect),
set_current_vllm_config(self.vllm_config),
):
self.proposer._propose(target_token_ids, target_positions, target_hidden_states, next_token_ids,
token_indices_to_sample, mock_common_attn_metadata, target_model_batch_desc, mock_sampling_metadata,
mm_embed_inputs, req_scheduled_tokens, long_seq_metadata, num_prefill_reqs, num_decode_reqs,
scheduler_output, num_scheduled_tokens, num_rejected_tokens_gpu,
)
self.assert_value_common_attn_metadata(captured_common_attn_metadata, flag_prefill_decode, model_type, graphmode)
# give common_attn_metadata value
def value_mock_common_attn_metadata(self, mock_common_attn_metadata, query_start_loc, query_start_loc_cpu, seq_lens, num_reqs,
num_actual_tokens, max_query_len, max_seq_len, block_table_tensor,
slot_mapping, causal, logits_indices_padded, num_logits_indices,
encoder_seq_lens, encoder_seq_lens_cpu, dcp_local_seq_lens,
dcp_local_seq_lens_cpu, _seq_lens_cpu, _num_computed_tokens_cpu,
_num_computed_tokens_cache, seq_lens_cpu, num_computed_tokens_cpu,
decode_token_per_req, actual_seq_lengths_q, positions, attn_state,
graph_pad_size, num_input_tokens, prefill_context_parallel_metadata
):
mock_common_attn_metadata.query_start_loc = query_start_loc
mock_common_attn_metadata.query_start_loc_cpu = query_start_loc_cpu
mock_common_attn_metadata.seq_lens = seq_lens
mock_common_attn_metadata.num_reqs = num_reqs
mock_common_attn_metadata.num_actual_tokens = num_actual_tokens
mock_common_attn_metadata.max_query_len = max_query_len
mock_common_attn_metadata.max_seq_len = max_seq_len
mock_common_attn_metadata.block_table_tensor = block_table_tensor
mock_common_attn_metadata.slot_mapping = slot_mapping
mock_common_attn_metadata.causal = causal
mock_common_attn_metadata.logits_indices_padded = logits_indices_padded
mock_common_attn_metadata.num_logits_indices = num_logits_indices
mock_common_attn_metadata.encoder_seq_lens = encoder_seq_lens
mock_common_attn_metadata.encoder_seq_lens_cpu = encoder_seq_lens_cpu
mock_common_attn_metadata.dcp_local_seq_lens = dcp_local_seq_lens
mock_common_attn_metadata.dcp_local_seq_lens_cpu = dcp_local_seq_lens_cpu
mock_common_attn_metadata._seq_lens_cpu = _seq_lens_cpu
mock_common_attn_metadata._num_computed_tokens_cpu = _num_computed_tokens_cpu
mock_common_attn_metadata._num_computed_tokens_cache = _num_computed_tokens_cache
mock_common_attn_metadata.seq_lens_cpu = seq_lens_cpu
mock_common_attn_metadata.num_computed_tokens_cpu = num_computed_tokens_cpu
mock_common_attn_metadata.decode_token_per_req = decode_token_per_req
mock_common_attn_metadata.actual_seq_lengths_q = actual_seq_lengths_q
mock_common_attn_metadata.positions = positions
mock_common_attn_metadata.attn_state = attn_state
mock_common_attn_metadata.graph_pad_size = graph_pad_size
mock_common_attn_metadata.num_input_tokens = num_input_tokens
mock_common_attn_metadata.prefill_context_parallel_metadata = prefill_context_parallel_metadata
# assert the value common_attn_metadata
def assert_value_common_attn_metadata(self, captured_common_attn_metadata, flag_prefill_decode, model_type, graphmode):
if not self.is_decode(flag_prefill_decode):
if flag_prefill_decode == 'prefill':
assert torch.equal(captured_common_attn_metadata.query_start_loc, torch.tensor([0, 1]))
assert torch.equal(captured_common_attn_metadata.query_start_loc_cpu, torch.tensor([0, 1]))
assert captured_common_attn_metadata.num_reqs == 1
assert captured_common_attn_metadata.num_actual_tokens == 1
assert captured_common_attn_metadata.max_query_len == 1
assert torch.equal(captured_common_attn_metadata.block_table_tensor, torch.eye(256, dtype=torch.int32)[0].unsqueeze(0))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([2]))
if flag_prefill_decode == 'decode_and_prefill':
assert torch.equal(captured_common_attn_metadata.query_start_loc, torch.tensor([0, 1, 2, 3]))
assert torch.equal(captured_common_attn_metadata.query_start_loc_cpu, torch.tensor([0, 1, 2, 3]))
assert captured_common_attn_metadata.num_reqs == 3
assert captured_common_attn_metadata.num_actual_tokens == 3
assert captured_common_attn_metadata.max_query_len == 1
assert torch.equal(captured_common_attn_metadata.block_table_tensor, torch.cat([torch.eye(256, device="cpu", dtype=torch.int32)[0].unsqueeze(0)*i for i in [1,2,3]], dim=0))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([15, 2, 2]))
if model_type == 'qwen_moe':
assert captured_common_attn_metadata._seq_lens_cpu == torch.tensor([15])
if model_type == 'qwen_dense':
if flag_prefill_decode == 'prefill':
assert captured_common_attn_metadata._seq_lens_cpu is None
if flag_prefill_decode == 'decode_and_prefill':
assert torch.equal(captured_common_attn_metadata._seq_lens_cpu, torch.tensor([19, 15, 15]))
if model_type == 'deepseek':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([11]))
assert captured_common_attn_metadata.max_seq_len == 9
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([138]), torch.full((9215,), -1)]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([11]))
assert captured_common_attn_metadata._seq_lens_cpu == torch.tensor([11])
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([10, 1, 2, 3, 4, 5, 6, 7, 8] + [0]*(8704-9), dtype=torch.int64))
assert captured_common_attn_metadata.num_input_tokens == 9
if model_type == 'qwen_dense' or model_type == 'qwen_moe':
if flag_prefill_decode == 'prefill':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([15]))
assert captured_common_attn_metadata.max_seq_len == 13
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([142]), torch.full((9215,), -1)]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([15]))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([14, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] + [0]*(8704-13), dtype=torch.int64))
assert captured_common_attn_metadata.num_input_tokens == 13
if flag_prefill_decode == 'decode_and_prefill':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([19, 15, 15]))
assert captured_common_attn_metadata.max_seq_len == 0
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([146, 270, 398]), torch.full((9213,), -1)]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([19, 15, 15]))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([18, 14, 14, 16, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12,
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] + [0]*(8704-30), dtype=torch.int64))
assert captured_common_attn_metadata.num_input_tokens == 30
if self.is_decode(flag_prefill_decode):
if graphmode == 'full':
assert torch.equal(captured_common_attn_metadata.query_start_loc, torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]))
assert torch.equal(captured_common_attn_metadata.query_start_loc_cpu, torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]))
assert captured_common_attn_metadata.num_reqs == 16
assert torch.equal(captured_common_attn_metadata.block_table_tensor, torch.cat([torch.eye(256, device="cpu", dtype=torch.int32)[0].unsqueeze(0)*i for i in [1,2,3]]
+ [torch.zeros(13, 256, device="cpu", dtype=torch.int32)], dim=0))
assert captured_common_attn_metadata.num_input_tokens == 16
else:
assert torch.equal(captured_common_attn_metadata.query_start_loc, torch.tensor([0, 1, 2, 3]))
assert torch.equal(captured_common_attn_metadata.query_start_loc_cpu, torch.tensor([0, 1, 2, 3]))
assert captured_common_attn_metadata.num_reqs == 3
assert torch.equal(captured_common_attn_metadata.block_table_tensor, torch.cat([torch.eye(256, device="cpu", dtype=torch.int32)[0].unsqueeze(0)*i for i in [1,2,3]], dim=0))
assert captured_common_attn_metadata.num_input_tokens == 12
assert captured_common_attn_metadata.num_actual_tokens == 3
assert captured_common_attn_metadata.max_query_len == 1
assert captured_common_attn_metadata.max_seq_len == 0
assert captured_common_attn_metadata._seq_lens_cpu is None
if model_type == 'qwen_dense':
if graphmode == 'full':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([23, 19, 19] + [0]*13))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([23, 19, 19] + [0]*13))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([19, 15, 15] + [0]*13))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([20, 18, 18, 20, 13, 14, 15, 16, 13, 14, 15, 16, 0, 0, 0, 0, 12, 0, 1,
2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] + [0]*(8704-30), dtype=torch.int64))
else:
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([23, 19, 19]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([23, 19, 19]))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([19, 15, 15]))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([20, 18, 18, 20, 13, 14, 15, 16, 13, 14, 15, 16, 8, 9, 10, 11, 12, 0, 1,
2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] + [0]*(8704-30), dtype=torch.int64))
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([148, 274, 402]), torch.full((9213,), -1)]))
if model_type == 'qwen_moe':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([21, 19, 19]))
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([145, 272, 400]), torch.full((9213,), -1)]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([21, 19, 19]))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([17, 15, 15]))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([17, 16, 16, 18, 13, 14, 15, 16, 13, 14, 15, 16, 8, 9, 10, 11, 12, 0, 1,
2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] + [0]*(8704-30), dtype=torch.int64))
if model_type == 'deepseek':
assert torch.equal(captured_common_attn_metadata.seq_lens, torch.tensor([16, 15, 16]))
assert torch.equal(captured_common_attn_metadata.slot_mapping, torch.cat([torch.tensor([142, 268, 396]), torch.full((9213,), -1)]))
assert torch.equal(captured_common_attn_metadata.seq_lens_cpu, torch.tensor([16, 15, 16]))
assert torch.equal(captured_common_attn_metadata.num_computed_tokens_cpu, torch.tensor([12, 11, 12]))
assert torch.equal(captured_common_attn_metadata.positions, torch.tensor([14, 12, 12, 13, 9, 10, 11, 12, 10, 11, 12, 13, 8, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9] + [0]*(8704-23), dtype=torch.int64))
assert captured_common_attn_metadata.causal
assert captured_common_attn_metadata.logits_indices_padded is None
assert captured_common_attn_metadata.num_logits_indices is None
assert captured_common_attn_metadata.encoder_seq_lens is None
assert captured_common_attn_metadata.encoder_seq_lens_cpu is None
assert captured_common_attn_metadata.dcp_local_seq_lens is None
assert captured_common_attn_metadata.dcp_local_seq_lens_cpu is None
assert captured_common_attn_metadata._num_computed_tokens_cpu is None
assert captured_common_attn_metadata._num_computed_tokens_cache is None
assert captured_common_attn_metadata.decode_token_per_req == 1
assert captured_common_attn_metadata.actual_seq_lengths_q == []
if model_type == 'deepseek':
assert captured_common_attn_metadata.attn_state == AscendAttentionState.SpecDecoding
else:
assert captured_common_attn_metadata.attn_state == AscendAttentionState.ChunkedPrefill
assert captured_common_attn_metadata.graph_pad_size == -1
assert captured_common_attn_metadata.prefill_context_parallel_metadata is None
if model_type == 'qwen_dense' and graphmode == 'eager' and flag_prefill_decode == 'decode_and_prefill':
assert torch.equal(self.proposer.slot_mapping_group[0][:30], torch.tensor([141, 142, 143, 144, 256, 257, 258, 259, 260, 261, 262, 263, 264, 265,
266, 267, 268, 384, 385, 386, 387, 388, 389, 390, 391, 392, 393, 394, 395, 396]))
assert torch.equal(self.proposer.slot_mapping_group[1][:3], torch.tensor([145, 269, 397]))
assert torch.equal(self.proposer.slot_mapping_group[2][:3], torch.tensor([146, 270, 398]))
assert self.proposer.slot_mapping_group[0].shape == torch.Size([9216])
assert torch.equal(self.proposer.seq_lens_group[0][:3], torch.tensor([17, 13, 13]))
assert torch.equal(self.proposer.seq_lens_group[1][:3], torch.tensor([18, 14, 14]))
assert torch.equal(self.proposer.seq_lens_group[2][:3], torch.tensor([19, 15, 15]))
assert self.proposer.seq_lens_group[0].shape == torch.Size([8704])
assert torch.equal(self.proposer.query_start_loc_group[0][:4], torch.tensor([0, 4, 17, 30]))
assert torch.equal(self.proposer.query_start_loc_group[1][:4], torch.tensor([0, 1, 2, 3]))
assert torch.equal(self.proposer.query_start_loc_group[2][:4], torch.tensor([0, 1, 2, 3]))
assert self.proposer.query_start_loc_group[0].shape == torch.Size([8704])
assert torch.equal(self.proposer.token_indices_to_sample[:3], torch.tensor([3, 16, 29]))
assert self.proposer.token_indices_to_sample.shape == torch.Size([1024])
# prefill or decode
def is_decode(self, flag_prefill_decode):
if flag_prefill_decode == "decode":
return True
elif flag_prefill_decode == "prefill":
return False
else:
return False
# Add assertions to ensure that the mocked functions and parameters exist
def check_mock(self):
import vllm.config
assert hasattr(vllm.config, "VllmConfig"), "VllmConfig not found"
fields = {
"speculative_config",
"scheduler_config",
"model_config",
"parallel_config",
"additional_config",
}
actual = set(vllm.config.VllmConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "SpeculativeConfig"), "SpeculativeConfig not found"
fields = {
"num_speculative_tokens",
"method",
"parallel_drafting",
"draft_tensor_parallel_size",
"draft_model_config",
"disable_padded_drafter_batch",
}
# speculative_token_tree was removed in newer vllm (Remove tree attention #42121);
# only check for it when the installed version still carries the field.
if "speculative_token_tree" in vllm.config.SpeculativeConfig.__dataclass_fields__:
fields.add("speculative_token_tree")
actual = set(vllm.config.SpeculativeConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "SchedulerConfig")
assert "max_num_batched_tokens" in vllm.config.SchedulerConfig.__dataclass_fields__
assert "max_num_seqs" in vllm.config.SchedulerConfig.__dataclass_fields__
assert hasattr(vllm.config, "ModelConfig")
assert "dtype" in vllm.config.ModelConfig.__dataclass_fields__
assert "max_model_len" in vllm.config.ModelConfig.__dataclass_fields__
assert isinstance(
inspect.getattr_static(vllm.config.ModelConfig, "uses_mrope"),
property
)
assert isinstance(
inspect.getattr_static(vllm.config.ModelConfig, "uses_xdrope_dim"),
property
)
assert isinstance(
inspect.getattr_static(vllm.config.ModelConfig, "use_mla"),
property
)
assert hasattr(vllm.config, "ParallelConfig"), "ParallelConfig not found"
fields = {
"tensor_parallel_size",
"data_parallel_rank",
"data_parallel_size",
"prefill_context_parallel_size",
}
actual = set(vllm.config.ParallelConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
import vllm_ascend.worker.model_runner_v1
assert hasattr(vllm_ascend.worker.model_runner_v1, "NPUModelRunner")
RunnerCls = vllm_ascend.worker.model_runner_v1.NPUModelRunner
src = inspect.getsource(RunnerCls.__init__)
fields = {
"pcp_size",
"dcp_size",
"max_num_tokens",
"max_num_reqs",
"pin_memory",
"query_start_loc",
}
for f in fields:
assert f"self.{f}" in src, f"missing self.{f} in __init__"
assert hasattr(RunnerCls, "_sync_metadata_across_dp")
sig = inspect.signature(RunnerCls._sync_metadata_across_dp)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'num_tokens', 'is_draft_model', 'cudagraph_mode', 'allow_dp_padding']
assert hasattr(RunnerCls, "_pad_query_start_loc_for_fia")
sig = inspect.signature(RunnerCls._pad_query_start_loc_for_fia)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'query_start_loc', 'num_tokens_padded', 'num_reqs_padded', 'num_reqs', 'cudagraph_runtime_mode', 'batch_desc_num_reqs']
import vllm_ascend.spec_decode.llm_base_proposer
assert hasattr(vllm_ascend.spec_decode.llm_base_proposer, "AscendSpecDecodeBaseProposer")
RunnerCls = vllm_ascend.spec_decode.llm_base_proposer.AscendSpecDecodeBaseProposer
assert hasattr(RunnerCls, "_get_model")
assert hasattr(RunnerCls, "_update_full_graph_params")
assert hasattr(RunnerCls, "_propose")
sig = inspect.signature(RunnerCls._get_model)
sig_name = self.get_param_names(sig)
assert sig_name == ['self']
sig = inspect.signature(RunnerCls._update_full_graph_params)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'forward_context', 'num_tokens', 'draft_attn_metadatas']
src = inspect.getsource(RunnerCls.load_model)
assert 'self.attn_layer_names' in src
assert 'self.kernel_block_size' in src
sig = inspect.signature(RunnerCls._propose)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'target_token_ids', 'target_positions', 'target_hidden_states', 'next_token_ids',
'token_indices_to_sample', 'common_attn_metadata', 'target_model_batch_desc',
'sampling_metadata', 'mm_embed_inputs', 'req_scheduled_tokens', 'long_seq_metadata',
'num_prefill_reqs', 'num_decode_reqs', 'scheduler_output', 'num_scheduled_tokens',
'num_rejected_tokens_gpu'
]
import vllm.model_executor.models.llama_eagle3
assert hasattr(vllm.model_executor.models.llama_eagle3, "Eagle3LlamaForCausalLM")
RunnerCls = vllm.model_executor.models.llama_eagle3.Eagle3LlamaForCausalLM
assert hasattr(RunnerCls, "combine_hidden_states")
sig = inspect.signature(RunnerCls.combine_hidden_states)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'hidden_states']
import vllm.v1.spec_decode.eagle
assert hasattr(vllm.v1.spec_decode.eagle, 'SpecDecodeBaseProposer')
RunnerCls = vllm.v1.spec_decode.eagle.SpecDecodeBaseProposer
src = inspect.getsource(RunnerCls.__init__)
assert 'self.hidden_size' in src
assert 'self.draft_attn_groups' in src
import vllm.v1.worker.gpu_model_runner
assert hasattr(vllm.v1.worker.gpu_model_runner, 'GPUModelRunner')
RunnerCls = vllm.v1.worker.gpu_model_runner.GPUModelRunner
src = inspect.getsource(RunnerCls.__init__)
assert 'self.cudagraph_dispatcher' in src
assert 'self.seq_lens' in src
assert 'self.optimistic_seq_lens_cpu' in src
import vllm.v1.cudagraph_dispatcher
assert hasattr(vllm.v1.cudagraph_dispatcher, 'CudagraphDispatcher')
assert hasattr(vllm.v1.cudagraph_dispatcher.CudagraphDispatcher, 'dispatch')
RunnerCls = vllm.v1.cudagraph_dispatcher.CudagraphDispatcher
sig = inspect.signature(RunnerCls.dispatch)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'num_tokens', 'uniform_decode', 'has_lora', 'num_active_loras', 'valid_modes', 'invalid_modes']
import vllm.v1.attention.backend
assert hasattr(vllm.v1.attention.backend, 'CommonAttentionMetadata')
fields = {
'query_start_loc', 'query_start_loc_cpu', 'seq_lens', 'num_reqs', \
'num_actual_tokens', 'max_query_len', 'max_seq_len', 'block_table_tensor', \
'slot_mapping', 'causal', 'logits_indices_padded', 'num_logits_indices', \
'encoder_seq_lens', 'encoder_seq_lens_cpu', 'dcp_local_seq_lens', \
'dcp_local_seq_lens_cpu', '_seq_lens_cpu', '_num_computed_tokens_cpu', \
'_num_computed_tokens_cache'
}
actual = set(vllm.v1.attention.backend.CommonAttentionMetadata.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
import vllm_ascend.attention.utils
assert hasattr(vllm_ascend.attention.utils, 'AscendCommonAttentionMetadata')
fields = {
'positions', 'seq_lens_cpu', 'decode_token_per_req', \
'prefill_context_parallel_metadata', 'actual_seq_lengths_q', \
'attn_state', 'num_computed_tokens_cpu', 'num_input_tokens', \
'graph_pad_size'
}
actual = set(vllm_ascend.attention.utils.AscendCommonAttentionMetadata.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
import vllm_ascend.spec_decode.llm_base_proposer
assert hasattr(vllm_ascend.spec_decode.llm_base_proposer, "AscendSpecDecodeBaseProposer")
RunnerCls = vllm_ascend.spec_decode.llm_base_proposer.AscendSpecDecodeBaseProposer
assert hasattr(RunnerCls, "_run_merged_draft")
sig = inspect.signature(RunnerCls._run_merged_draft)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'num_input_tokens', 'batch_size', 'token_indices_to_sample',
'target_positions', 'inputs_embeds', 'multi_steps_attn_metadata',
'num_tokens', 'is_prefill'
]
import vllm.v1.worker.utils
assert hasattr(vllm.v1.worker.utils, "AttentionGroup")
assert hasattr(vllm.v1.worker.utils.AttentionGroup, "get_metadata_builder")
fields = {
'backend', 'layer_names', 'kv_cache_spec', \
'kv_cache_group_id'
}
actual = set(vllm.v1.worker.utils.AttentionGroup.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
import vllm.v1.attention.backend
assert hasattr(vllm.v1.attention.backend, "AttentionMetadataBuilder")
assert hasattr(vllm.v1.attention.backend.AttentionMetadataBuilder, "build")
assert hasattr(vllm.v1.attention.backend.AttentionMetadataBuilder, "build_for_drafting")
RunnerCls = vllm.v1.attention.backend.AttentionMetadataBuilder
sig = inspect.signature(RunnerCls.build)
sig_name = self.get_param_names(sig)
assert sig_name == ['self', 'common_prefix_len', 'common_attn_metadata', 'fast_build']
# get the param in inspect sig, for check_mock()
def get_param_names(self, sig):
return [p.name for p in sig.parameters.values()]
# fmt: on
class MockCpuGpuBuffer:
"""Mock CpuGpuBuffer for testing"""
def __init__(self, max_size, dtype, device="cpu", **kwargs):
self.max_size = max_size
self.dtype = dtype
self.device = device
self.cpu = torch.zeros(max_size, dtype=dtype, device="cpu")
self.np = self.cpu.numpy()
self.gpu = torch.zeros(max_size, dtype=dtype, device=device)
def copy_to_gpu(self, size=None):
if size is None:
size = self.max_size
self.gpu[:size].copy_(self.cpu[:size])
class MockCachedRequestState:
"""Mock CachedRequestState for testing"""
def __init__(self, req_id, token_ids):
self.req_id = req_id
self.token_ids = token_ids
def get_token_id(self, position):
if position < len(self.token_ids):
return self.token_ids[position]
return 0
class MockInputBatch:
"""Mock InputBatch for testing"""
def __init__(self, num_reqs, req_ids, vocab_size, num_tokens_no_spec=None):
self.num_reqs = num_reqs
self.req_ids = req_ids
self.vocab_size = vocab_size
# num_tokens_no_spec represents the sequence length (excluding speculative tokens)
# for each request. Default to seq_len + 1 for each request.
if num_tokens_no_spec is None:
self.num_tokens_no_spec = np.array([i + 11 for i in range(num_reqs)], dtype=np.int64)
else:
self.num_tokens_no_spec = np.array(num_tokens_no_spec, dtype=np.int64)
class TestPrepareNextTokenIdsPadded(TestBase):
"""Test prepare_next_token_ids_padded method with precision validation"""
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.cache_config = MagicMock(spec=CacheConfig)
self.vllm_config.scheduler_config = MagicMock()
self.vllm_config.model_config = MagicMock()
self.vllm_config.model_config.hf_text_config = MagicMock(spec=[])
self.vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
self.vllm_config.compilation_config = MagicMock()
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.num_speculative_tokens = 4
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(4)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
def tearDown(self):
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
set_current_vllm_config(None)
def test_all_valid_tokens(self):
"""Test case where all requests have valid sampled tokens"""
num_reqs = 3
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[100, 101, 102, 103, 104],
[200, 201, 202, 203, 204],
[300, 301, 302, 303, 304],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(10))),
"req_1": MockCachedRequestState("req_1", list(range(15))),
"req_2": MockCachedRequestState("req_2", list(range(20))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1", "req_2"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 16, 21], # seq_len = num_tokens_no_spec - 1 = [10, 15, 20]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
self.assertEqual(next_token_ids.shape[0], num_reqs)
self.assertEqual(valid_sampled_tokens_count.shape[0], num_reqs)
expected_valid_counts = torch.tensor([5, 5, 5], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_next_tokens = torch.tensor([104, 204, 304], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_partial_rejected_tokens(self):
"""Test case where some tokens are rejected (marked as -1)"""
num_reqs = 3
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[100, 101, -1, -1, -1],
[200, 201, 202, 203, -1],
[300, 301, 302, 303, 304],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(10))),
"req_1": MockCachedRequestState("req_1", list(range(15))),
"req_2": MockCachedRequestState("req_2", list(range(20))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1", "req_2"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 16, 21], # seq_len = num_tokens_no_spec - 1 = [10, 15, 20]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
expected_valid_counts = torch.tensor([2, 4, 5], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_next_tokens = torch.tensor([101, 203, 304], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_all_rejected_tokens_with_backup(self):
"""Test case where all tokens are rejected, should use backup token"""
num_reqs = 3
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[-1, -1, -1, -1, -1],
[-1, -1, -1, -1, -1],
[300, 301, 302, 303, 304],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
"req_1": MockCachedRequestState("req_1", list(range(20))),
"req_2": MockCachedRequestState("req_2", list(range(25))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1", "req_2"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 16, 26], # seq_len = num_tokens_no_spec - 1 = [10, 15, 25]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
expected_valid_counts = torch.tensor([0, 0, 5], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_backup_token_0 = requests["req_0"].get_token_id(10)
expected_backup_token_1 = requests["req_1"].get_token_id(15)
expected_next_tokens = torch.tensor([expected_backup_token_0, expected_backup_token_1, 304], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_discarded_requests(self):
"""Test case with discarded requests"""
num_reqs = 3
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[100, 101, 102, 103, 104],
[200, 201, 202, 203, 204],
[300, 301, 302, 303, 304],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
"req_1": MockCachedRequestState("req_1", list(range(20))),
"req_2": MockCachedRequestState("req_2", list(range(25))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1", "req_2"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 21, 21], # seq_len = num_tokens_no_spec - 1 = [10, 20, 20]
)
discard_request_indices = torch.tensor([0, 2], dtype=torch.int64)
num_discarded_requests = 2
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
expected_valid_counts = torch.tensor([0, 5, 0], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_backup_token_0 = requests["req_0"].get_token_id(10)
expected_backup_token_2 = requests["req_2"].get_token_id(20)
expected_next_tokens = torch.tensor([expected_backup_token_0, 204, expected_backup_token_2], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_mixed_scenario(self):
"""Test mixed scenario: some rejected tokens, some discarded requests, some all-rejected"""
num_reqs = 4
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[100, 101, -1, -1, -1],
[-1, -1, -1, -1, -1],
[300, 301, 302, 303, 304],
[400, 401, 402, -1, -1],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
"req_1": MockCachedRequestState("req_1", list(range(20))),
"req_2": MockCachedRequestState("req_2", list(range(25))),
"req_3": MockCachedRequestState("req_3", list(range(30))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1", "req_2", "req_3"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 16, 26, 31], # seq_len = num_tokens_no_spec - 1 = [10, 15, 25, 30]
)
discard_request_indices = torch.tensor([1], dtype=torch.int64)
num_discarded_requests = 1
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
expected_valid_counts = torch.tensor([2, 0, 5, 3], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_backup_token_1 = requests["req_1"].get_token_id(15)
expected_next_tokens = torch.tensor([101, expected_backup_token_1, 304, 402], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_single_request(self):
"""Test with single request"""
num_reqs = 1
vocab_size = 1000
sampled_token_ids = torch.tensor([[100, 101, 102, 103, 104]], dtype=torch.int64)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0"],
vocab_size=vocab_size,
num_tokens_no_spec=[16], # seq_len = num_tokens_no_spec - 1 = [15]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
expected_valid_counts = torch.tensor([5], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
expected_next_tokens = torch.tensor([104], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_vocab_size_boundary(self):
"""Test with tokens at vocab size boundary"""
num_reqs = 2
vocab_size = 100
sampled_token_ids = torch.tensor(
[
[99, 100, 101, -1, -1],
[50, 51, 52, 53, 54],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
"req_1": MockCachedRequestState("req_1", list(range(20))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1"],
vocab_size=vocab_size,
num_tokens_no_spec=[16, 21], # seq_len = num_tokens_no_spec - 1 = [15, 20]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
# Token 100 and 101 are >= vocab_size (100), so they are invalid
# Only token 99 is valid for the first request
expected_valid_counts = torch.tensor([1, 5], dtype=torch.int64)
self.assertTrue(torch.equal(valid_sampled_tokens_count, expected_valid_counts))
# Next token should be 99 (last valid token) for first request
# and 54 for second request
expected_next_tokens = torch.tensor([99, 54], dtype=torch.int64)
self.assertTrue(torch.equal(next_token_ids, expected_next_tokens))
def test_intermediate_variables_precision(self):
"""Test to verify key variables that affect downstream computation"""
num_reqs = 2
vocab_size = 1000
sampled_token_ids = torch.tensor(
[
[100, 101, -1, -1, -1],
[200, 201, 202, -1, -1],
],
dtype=torch.int64,
)
requests = {
"req_0": MockCachedRequestState("req_0", list(range(15))),
"req_1": MockCachedRequestState("req_1", list(range(20))),
}
gpu_input_batch = MockInputBatch(
num_reqs=num_reqs,
req_ids=["req_0", "req_1"],
vocab_size=vocab_size,
num_tokens_no_spec=[11, 16], # seq_len = num_tokens_no_spec - 1 = [10, 15]
)
discard_request_indices = torch.tensor([], dtype=torch.int64)
num_discarded_requests = 0
next_token_ids, valid_sampled_tokens_count = self.proposer.prepare_next_token_ids_padded(
sampled_token_ids=sampled_token_ids,
requests=requests,
gpu_input_batch=gpu_input_batch,
discard_request_indices=discard_request_indices,
num_discarded_requests=num_discarded_requests,
)
# Verify return values
self.assertEqual(next_token_ids[0].item(), 101, "next_token_ids[0] should be 101")
self.assertEqual(next_token_ids[1].item(), 202, "next_token_ids[1] should be 202")
self.assertEqual(valid_sampled_tokens_count[0].item(), 2, "valid_sampled_tokens_count[0] should be 2")
self.assertEqual(valid_sampled_tokens_count[1].item(), 3, "valid_sampled_tokens_count[1] should be 3")
# Verify public member that affects downstream computation
expected_backup_0 = requests["req_0"].get_token_id(10)
expected_backup_1 = requests["req_1"].get_token_id(15)
self.assertEqual(
self.proposer.backup_next_token_ids.np[0],
expected_backup_0,
f"backup_next_token_ids[0] should be {expected_backup_0}",
)
self.assertEqual(
self.proposer.backup_next_token_ids.np[1],
expected_backup_1,
f"backup_next_token_ids[1] should be {expected_backup_1}",
)
# Verify data types
self.assertEqual(next_token_ids.dtype, torch.int64, "next_token_ids dtype should be torch.int64")
self.assertEqual(
valid_sampled_tokens_count.dtype, torch.int64, "valid_sampled_tokens_count dtype should be torch.int64"
)
# fmt: off
class MockDraftModel:
"""Draft model that records prepared forward inputs."""
def __init__(self, returns_tuple=True, vocab_size=200000):
self.returns_tuple = returns_tuple
self.vocab_size = vocab_size
self.calls = []
self.logit_inputs = []
self.returned_hidden_states = []
def __call__(self, **kwargs):
self.calls.append({key: value.clone() if torch.is_tensor(value) else value for key, value in kwargs.items()})
input_ids = kwargs["input_ids"].to(torch.long)
call_idx = len(self.returned_hidden_states)
last_hidden_states = torch.zeros(input_ids.shape[0], 4, dtype=torch.float32)
last_hidden_states[:, 0] = input_ids + call_idx
last_hidden_states[:, 1] = 100 + input_ids + call_idx
hidden_states = torch.zeros_like(last_hidden_states)
hidden_states[:, 0] = 1000 + input_ids + call_idx
hidden_states[:, 1] = 2000 + input_ids + call_idx
self.returned_hidden_states.append((last_hidden_states.clone(), hidden_states.clone()))
if self.returns_tuple:
return last_hidden_states, hidden_states
return last_hidden_states
def compute_logits(self, sample_hidden_states):
self.logit_inputs.append(sample_hidden_states.clone())
token_ids = sample_hidden_states[:, 0].to(torch.long)
logits = torch.full((sample_hidden_states.shape[0], self.vocab_size), -1000.0)
logits[torch.arange(sample_hidden_states.shape[0]), token_ids] = 1000.0
return logits
def embed_input_ids(self, input_ids):
return torch.stack((input_ids.float() + 5000, input_ids.float() + 6000), dim=1).repeat(1, 2)
class TestRunMergedDraft(TestBase):
def setUp(self):
self.check_mock()
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.cache_config = MagicMock(spec=CacheConfig)
self.vllm_config.scheduler_config = MagicMock()
self.vllm_config.model_config = MagicMock()
self.vllm_config.model_config.hf_text_config = MagicMock(spec=[])
self.vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
self.vllm_config.compilation_config = MagicMock()
self.vllm_config.compilation_config.mode = CompilationMode.NONE
self.vllm_config.compilation_config.pass_config = MagicMock()
self.vllm_config.compilation_config.pass_config.enable_sp = False
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_rank = 0
self.runner.dcp_rank = 0
self.runner.pcp_manager = None
self.runner.max_num_tokens = 64
self.runner.max_num_reqs = 8
self.runner.uniform_decode_query_len = 2
self.runner.enable_enpu = False
self.runner.use_eagle = True
self.runner._use_aclgraph.return_value = False
self.runner._make_buffer.side_effect = lambda size, dtype: torch.zeros(size, dtype=dtype, device=self.device)
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.async_scheduling = False
self.vllm_config.scheduler_config.max_num_batched_tokens = 64
self.vllm_config.scheduler_config.max_num_seqs = 8
self.vllm_config.model_config.dtype = torch.float32
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.model_config.use_mla = False
self.vllm_config.model_config.is_multimodal_model = False
self.vllm_config.model_config.enforce_eager = True
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.method = "eagle3"
self.vllm_config.speculative_config.num_speculative_tokens = 3
self.vllm_config.speculative_config.parallel_drafting = False
self.vllm_config.speculative_config.enforce_eager = True
self.vllm_config.speculative_config.use_local_argmax_reduction = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(3)])
self.vllm_config.speculative_config.draft_model_config.hf_config = MagicMock(spec=[])
self.vllm_config.speculative_config.draft_model_config.get_hidden_size.return_value = 4
self.vllm_config.speculative_config.draft_model_config.get_inputs_embeds_size.return_value = 4
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
self.mock_enable_sp = patch("vllm_ascend.spec_decode.llm_base_proposer.enable_sp", return_value=False)
self.mock_enable_sp.start()
self.mock_shared_expert_dp = patch(
"vllm_ascend.spec_decode.llm_base_proposer.shared_expert_dp_enabled", return_value=False
)
self.mock_shared_expert_dp.start()
self.mock_extra_ctx = patch("vllm_ascend.spec_decode.llm_base_proposer._EXTRA_CTX", new=MagicMock())
self.mock_extra_ctx.start()
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.proposer.maybe_pad_and_reduce = MagicMock(
side_effect=lambda hidden_states, positions: (hidden_states, positions)
)
self.proposer.maybe_all_gather_and_unpad = MagicMock(
side_effect=lambda last_hidden_states, positions, hidden_states: (
last_hidden_states,
positions,
hidden_states,
)
)
def tearDown(self):
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
self.mock_enable_sp.stop()
self.mock_shared_expert_dp.stop()
self.mock_extra_ctx.stop()
set_current_vllm_config(None)
def check_mock(self):
import vllm.config
assert hasattr(vllm.config, "VllmConfig"), "VllmConfig not found"
fields = {
"speculative_config",
"cache_config",
"scheduler_config",
"model_config",
"parallel_config",
"compilation_config",
"additional_config",
}
actual = set(vllm.config.VllmConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "CacheConfig"), "CacheConfig not found"
assert "block_size" in vllm.config.CacheConfig.__dataclass_fields__
assert hasattr(vllm.config, "SchedulerConfig"), "SchedulerConfig not found"
fields = {"async_scheduling", "max_num_batched_tokens", "max_num_seqs"}
actual = set(vllm.config.SchedulerConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "ModelConfig"), "ModelConfig not found"
fields = {"dtype", "max_model_len", "enforce_eager", "hf_text_config"}
actual = set(vllm.config.ModelConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
for field in ("uses_mrope", "uses_xdrope_dim", "use_mla", "is_multimodal_model"):
assert isinstance(inspect.getattr_static(vllm.config.ModelConfig, field), property)
for method in ("get_hidden_size", "get_inputs_embeds_size"):
assert hasattr(vllm.config.ModelConfig, method)
assert hasattr(vllm.config, "ParallelConfig"), "ParallelConfig not found"
fields = {
"tensor_parallel_size",
"data_parallel_rank",
"data_parallel_size",
"prefill_context_parallel_size",
"enable_expert_parallel",
}
actual = set(vllm.config.ParallelConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "SpeculativeConfig"), "SpeculativeConfig not found"
fields = {
"method",
"num_speculative_tokens",
"parallel_drafting",
"enforce_eager",
"use_local_argmax_reduction",
"draft_tensor_parallel_size",
"draft_model_config",
"disable_padded_drafter_batch",
}
# speculative_token_tree was removed in newer vllm (Remove tree attention #42121);
# only check for it when the installed version still carries the field.
if "speculative_token_tree" in vllm.config.SpeculativeConfig.__dataclass_fields__:
fields.add("speculative_token_tree")
actual = set(vllm.config.SpeculativeConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "CompilationConfig"), "CompilationConfig not found"
fields = {"mode", "pass_config"}
actual = set(vllm.config.CompilationConfig.__dataclass_fields__)
missing = fields - actual
assert not missing, f"Missing dataclass fields: {missing}"
assert hasattr(vllm.config, "PassConfig"), "PassConfig not found"
assert "enable_sp" in vllm.config.PassConfig.__dataclass_fields__
import vllm.forward_context
assert hasattr(vllm.forward_context, "get_forward_context")
import vllm.multimodal.registry
assert hasattr(vllm.multimodal.registry, "MultiModalRegistry")
assert hasattr(vllm.multimodal.registry.MultiModalRegistry, "supports_multimodal_inputs")
sig = inspect.signature(vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "model_config"]
import vllm.v1.spec_decode.eagle
# `CpuGpuBuffer` was re-exported from `eagle` until vLLM #40732 moved
# `SpecDecodeBaseProposer` (and the import) into `llm_base_proposer`.
import vllm.v1.spec_decode.llm_base_proposer
assert hasattr(vllm.v1.spec_decode.llm_base_proposer, "CpuGpuBuffer")
RunnerCls = vllm.v1.spec_decode.eagle.SpecDecodeBaseProposer
for attr in ("_get_positions", "_set_positions"):
assert hasattr(RunnerCls, attr), f"SpecDecodeBaseProposer.{attr} not found"
sig = inspect.signature(RunnerCls._get_positions)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "num_tokens"]
sig = inspect.signature(RunnerCls._set_positions)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "num_tokens", "positions"]
import vllm.model_executor.models.llama_eagle3
assert hasattr(vllm.model_executor.models.llama_eagle3, "Eagle3LlamaForCausalLM")
RunnerCls = vllm.model_executor.models.llama_eagle3.Eagle3LlamaForCausalLM
for attr in ("forward", "compute_logits", "embed_input_ids"):
assert hasattr(RunnerCls, attr), f"Eagle3LlamaForCausalLM.{attr} not found"
sig = inspect.signature(RunnerCls.forward)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "input_ids", "positions", "hidden_states", "inputs_embeds"]
sig = inspect.signature(RunnerCls.compute_logits)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "hidden_states"]
import vllm_ascend.ascend_forward_context
assert hasattr(vllm_ascend.ascend_forward_context, "_EXTRA_CTX")
extra_attrs = set(vllm_ascend.ascend_forward_context._ExtraForwardContextProxy.extra_attrs)
fields = {"num_tokens", "num_accept_tokens", "flash_comm_v1_enabled"}
missing = fields - extra_attrs
assert not missing, f"Missing extra forward context attrs: {missing}"
import vllm_ascend.spec_decode.llm_base_proposer
for attr in (
"AscendSpecDecodeBaseProposer",
"enable_sp",
"shared_expert_dp_enabled",
"lmhead_tp_enable",
"get_forward_context",
"_EXTRA_CTX",
):
assert hasattr(vllm_ascend.spec_decode.llm_base_proposer, attr), (
f"vllm_ascend.spec_decode.llm_base_proposer.{attr} not found"
)
RunnerCls = vllm_ascend.spec_decode.llm_base_proposer.AscendSpecDecodeBaseProposer
for attr in (
"_run_merged_draft",
"maybe_pad_and_reduce",
"maybe_all_gather_and_unpad",
"model_returns_tuple",
):
assert hasattr(RunnerCls, attr), f"AscendSpecDecodeBaseProposer.{attr} not found"
sig = inspect.signature(RunnerCls._run_merged_draft)
sig_name = self.get_param_names(sig)
assert sig_name == [
"self",
"num_input_tokens",
"batch_size",
"token_indices_to_sample",
"target_positions",
"inputs_embeds",
"multi_steps_attn_metadata",
"num_tokens",
"is_prefill",
]
sig = inspect.signature(RunnerCls.maybe_pad_and_reduce)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "hidden_states", "positions"]
sig = inspect.signature(RunnerCls.maybe_all_gather_and_unpad)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "last_hidden_states", "positions", "hidden_states"]
sig = inspect.signature(RunnerCls.model_returns_tuple)
sig_name = self.get_param_names(sig)
assert sig_name == ["self"]
import vllm_ascend.spec_decode.dflash_proposer
assert hasattr(vllm_ascend.spec_decode.dflash_proposer, "AscendDflashProposer")
RunnerCls = vllm_ascend.spec_decode.dflash_proposer.AscendDflashProposer
assert hasattr(RunnerCls, "build_model_inputs_first_pass")
sig = inspect.signature(RunnerCls.build_model_inputs_first_pass)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "num_input_tokens"]
import vllm_ascend.worker.model_runner_v1
assert hasattr(vllm_ascend.worker.model_runner_v1, "NPUModelRunner")
RunnerCls = vllm_ascend.worker.model_runner_v1.NPUModelRunner
src = inspect.getsource(RunnerCls.__init__)
fields = {
"pcp_size",
"dcp_size",
"pcp_rank",
"dcp_rank",
"max_num_tokens",
"max_num_reqs",
"uniform_decode_query_len",
"enable_enpu",
"use_eagle",
"pin_memory",
}
for f in fields:
assert f"self.{f}" in src, f"missing self.{f} in __init__"
assert hasattr(RunnerCls, "_use_aclgraph")
sig = inspect.signature(RunnerCls._use_aclgraph)
sig_name = self.get_param_names(sig)
assert sig_name == ["self"]
assert hasattr(RunnerCls, "_make_buffer")
sig = inspect.signature(RunnerCls._make_buffer)
sig_name = self.get_param_names(sig)
assert sig_name == ["self", "size", "dtype", "numpy"]
def get_param_names(self, sig):
return [p.name for p in sig.parameters.values()]
def test_run_merged_draft_eagle3_decode_prepares_each_forward_input(self):
self.proposer.model = MockDraftModel(returns_tuple=True)
def compute_draft_token_ids(sample_hidden_states):
self.proposer.model.logit_inputs.append(sample_hidden_states.clone())
token_ids = sample_hidden_states[:, 0].to(torch.long)
logits = torch.full((sample_hidden_states.shape[0], self.proposer.model.vocab_size), -1000.0)
logits[torch.arange(sample_hidden_states.shape[0]), token_ids] = 1000.0
logits = logits.argmax(dim=-1)
return logits
self.proposer.compute_draft_token_ids = compute_draft_token_ids
self.proposer.supports_mm_inputs = True
initial_input_ids = torch.tensor(
[279, 1196, 374, 8014, 151667, 198, 32313, 11, 151667, 198, 32313, 11],
dtype=torch.int32,
)
initial_positions = torch.tensor(
[17, 18, 19, 20, 13, 14, 15, 16, 13, 14, 15, 16],
dtype=torch.int32,
)
initial_hidden_states = torch.arange(48, dtype=torch.float32).view(12, 4)
self.proposer.input_ids[:12] = initial_input_ids
self.proposer.positions[:12] = initial_positions
self.proposer.hidden_states[:12] = initial_hidden_states
token_indices_to_sample = torch.tensor([1, 7, 11], dtype=torch.int64)
forward_context = MagicMock()
forward_context.moe_layer_index = 5
forward_context.attn_metadata = None
multi_steps_attn_metadata = [MagicMock(), MagicMock(), MagicMock()]
mock_ascend_config = MagicMock()
mock_ascend_config.enable_reduce_sample = True
with (
patch.object(llm_base_proposer, "lmhead_tp_enable", return_value=False),
patch.object(llm_base_proposer, "get_ascend_config", return_value=mock_ascend_config),
patch.object(llm_base_proposer, "get_forward_context", return_value=forward_context),
):
draft_token_ids = self.proposer._run_merged_draft(
num_input_tokens=12,
batch_size=3,
token_indices_to_sample=token_indices_to_sample,
target_positions=self.proposer.positions[:12],
inputs_embeds=None,
multi_steps_attn_metadata=multi_steps_attn_metadata,
num_tokens=12,
is_prefill=False,
)
model = self.proposer.model
self.assertEqual(draft_token_ids.tolist(), [[1196, 1197, 1199], [11, 12, 14], [11, 12, 14]])
self.assertEqual(len(model.calls), 3)
first_call = model.calls[0]
self.assertTrue(torch.equal(first_call["input_ids"], initial_input_ids))
self.assertTrue(torch.equal(first_call["positions"], initial_positions))
self.assertTrue(torch.equal(first_call["hidden_states"], initial_hidden_states))
self.assertIsNone(first_call["inputs_embeds"])
second_call = model.calls[1]
self.assertTrue(torch.equal(second_call["input_ids"], torch.tensor([1196, 11, 11], dtype=torch.int32)))
self.assertTrue(torch.equal(second_call["positions"], torch.tensor([19, 17, 17], dtype=torch.int32)))
self.assertTrue(
torch.equal(
second_call["hidden_states"],
model.returned_hidden_states[0][1][token_indices_to_sample],
)
)
self.assertTrue(
torch.equal(
second_call["inputs_embeds"],
model.embed_input_ids(torch.tensor([1196, 11, 11], dtype=torch.int32)),
)
)
third_call = model.calls[2]
self.assertTrue(torch.equal(third_call["input_ids"], torch.tensor([1197, 12, 12], dtype=torch.int32)))
self.assertTrue(torch.equal(third_call["positions"], torch.tensor([20, 18, 18], dtype=torch.int32)))
self.assertTrue(torch.equal(model.logit_inputs[0], model.returned_hidden_states[0][0][token_indices_to_sample]))
self.assertEqual(forward_context.moe_layer_index, 0)
self.assertIs(forward_context.attn_metadata, multi_steps_attn_metadata[2])
self.assertEqual(llm_base_proposer._EXTRA_CTX.num_tokens, 3)
self.assertEqual(llm_base_proposer._EXTRA_CTX.num_accept_tokens, 3)
def test_run_merged_draft_dflash_uses_first_pass_inputs_and_returns_early(self):
self.proposer.method = "dflash"
self.proposer.num_speculative_tokens = 1
self.proposer.pass_hidden_states_to_model = False
self.proposer.model = MockDraftModel(returns_tuple=False)
self.proposer.build_model_inputs_first_pass = MagicMock(
return_value={
"input_ids": torch.tensor([151667, 32313], dtype=torch.int32),
"positions": torch.tensor([20, 16], dtype=torch.int64),
"inputs_embeds": torch.ones(2, 4, dtype=torch.float32),
}
)
mock_ascend_config = MagicMock()
mock_ascend_config.enable_reduce_sample = False
with (
patch.object(llm_base_proposer, "lmhead_tp_enable", return_value=False),
patch.object(llm_base_proposer, "get_ascend_config", return_value=mock_ascend_config),
):
draft_token_ids = self.proposer._run_merged_draft(
num_input_tokens=12,
batch_size=2,
token_indices_to_sample=torch.tensor([0, 1], dtype=torch.int64),
target_positions=torch.tensor([20, 16], dtype=torch.int64),
inputs_embeds=None,
multi_steps_attn_metadata=None,
num_tokens=12,
is_prefill=False,
)
self.proposer.build_model_inputs_first_pass.assert_called_once_with(12)
self.proposer.maybe_all_gather_and_unpad.assert_not_called()
self.assertNotIn("hidden_states", self.proposer.model.calls[0])
self.assertTrue(
torch.equal(
self.proposer.model.calls[0]["input_ids"],
torch.tensor([151667, 32313], dtype=torch.int32),
)
)
self.assertEqual(draft_token_ids.tolist(), [[151667], [32313]])
def test_run_merged_draft_mtp_mrope_graph_and_lmhead_tp_preparation(self):
self.proposer.method = "mtp"
self.proposer.uses_mrope = True
self.proposer.use_cuda_graph = True
self.proposer.vllm_config.model_config.max_model_len = 4
self.proposer.vllm_config.scheduler_config.max_num_seqs = 3
self.proposer.runner.uniform_decode_query_len = 2
self.proposer.mrope_positions = torch.zeros((3, self.proposer.max_num_tokens + 1), dtype=torch.int64)
self.proposer.model = MockDraftModel(returns_tuple=False)
self.proposer.input_ids[:6] = torch.tensor([201, 33001, 14, 832, 128798, 271], dtype=torch.int32)
initial_mrope_positions = torch.tensor(
[
[0, 1, 3, 0, 1, 3],
[10, 11, 13, 10, 11, 13],
[20, 21, 23, 20, 21, 23],
],
dtype=torch.int64,
)
initial_hidden_states = torch.arange(24, dtype=torch.float32).view(6, 4)
self.proposer.mrope_positions[:, :6] = initial_mrope_positions
self.proposer.hidden_states[:6] = initial_hidden_states
token_indices_to_sample = torch.tensor([2, 5], dtype=torch.int64)
forward_context = MagicMock()
forward_context.moe_layer_index = 9
forward_context.attn_metadata = None
multi_steps_attn_metadata = [MagicMock(), MagicMock(), MagicMock()]
mock_ascend_config = MagicMock()
mock_ascend_config.enable_reduce_sample = False
with (
patch.object(llm_base_proposer, "lmhead_tp_enable", return_value=True),
patch.object(llm_base_proposer, "get_ascend_config", return_value=mock_ascend_config),
patch.object(llm_base_proposer, "get_forward_context", return_value=forward_context),
):
draft_token_ids = self.proposer._run_merged_draft(
num_input_tokens=6,
batch_size=2,
token_indices_to_sample=token_indices_to_sample,
target_positions=self.proposer.mrope_positions[:, :6],
inputs_embeds=None,
multi_steps_attn_metadata=multi_steps_attn_metadata,
num_tokens=6,
is_prefill=False,
)
model = self.proposer.model
self.assertEqual(draft_token_ids.tolist(), [[14, 15, 17], [271, 272, 274]])
self.assertTrue(all(logit_input.shape[0] == 6 for logit_input in model.logit_inputs))
first_call = model.calls[0]
self.assertTrue(torch.equal(first_call["positions"], initial_mrope_positions))
self.assertTrue(torch.equal(first_call["hidden_states"], initial_hidden_states))
second_call = model.calls[1]
self.assertTrue(
torch.equal(
second_call["input_ids"],
torch.tensor([14, 271, 14, 832, 128798, 271], dtype=torch.int32),
)
)
self.assertTrue(
torch.equal(
second_call["positions"],
torch.tensor(
[
[0, 0, 3, 0, 1, 3],
[0, 0, 13, 10, 11, 13],
[0, 0, 23, 20, 21, 23],
],
dtype=torch.int64,
),
)
)
self.assertTrue(
torch.equal(
second_call["hidden_states"][:2],
model.returned_hidden_states[0][0][token_indices_to_sample],
)
)
self.assertIs(forward_context.attn_metadata, multi_steps_attn_metadata[2])
def test_run_merged_draft_early_return_conditions(self):
test_cases = [
(1, False, torch.tensor([1, 3], dtype=torch.int64), (2, 1)),
(2, True, torch.tensor([0, 1, 2, 3], dtype=torch.int64), (2, 2)),
]
mock_ascend_config = MagicMock()
mock_ascend_config.enable_reduce_sample = False
for num_speculative_tokens, parallel_drafting, token_indices_to_sample, expected_shape in test_cases:
with self.subTest(num_speculative_tokens=num_speculative_tokens, parallel_drafting=parallel_drafting):
self.proposer.method = "eagle3"
self.proposer.num_speculative_tokens = num_speculative_tokens
self.proposer.parallel_drafting = parallel_drafting
self.proposer.pass_hidden_states_to_model = False
self.proposer.model = MockDraftModel(returns_tuple=True)
self.proposer.input_ids[:4] = torch.tensor([279, 1196, 374, 8014], dtype=torch.int32)
self.proposer.positions[:4] = torch.tensor([17, 18, 19, 20], dtype=torch.int64)
with (
patch.object(llm_base_proposer, "lmhead_tp_enable", return_value=False),
patch.object(llm_base_proposer, "get_ascend_config", return_value=mock_ascend_config),
):
draft_token_ids = self.proposer._run_merged_draft(
num_input_tokens=4,
batch_size=2,
token_indices_to_sample=token_indices_to_sample,
target_positions=self.proposer.positions[:4],
inputs_embeds=None,
multi_steps_attn_metadata=None,
num_tokens=4,
is_prefill=False,
)
self.assertEqual(tuple(draft_token_ids.shape), expected_shape)
self.assertEqual(len(self.proposer.model.calls), 1)
class TestDraftProposerHelperMethods(TestBase):
def setUp(self):
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.scheduler_config = MagicMock(max_num_seqs=3)
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.input_batch = MagicMock()
self.runner.input_batch.req_ids = [0, 1, 2]
self.runner.arange_np = np.arange(10)
self.runner.input_batch.num_reqs = 3
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.pcp_manager = None
self.vllm_config.cache_config.block_size = 16
self.vllm_config.scheduler_config.max_num_batched_tokens = 1024
self.vllm_config.scheduler_config.max_num_seqs = 32
self.vllm_config.model_config.dtype = torch.float16
self.vllm_config.model_config.max_model_len = 2048
self.vllm_config.model_config.uses_mrope = False
self.vllm_config.model_config.uses_xdrope_dim = 0
self.vllm_config.parallel_config.tensor_parallel_size = 1
self.vllm_config.parallel_config.data_parallel_rank = 0
self.vllm_config.parallel_config.data_parallel_size = 1
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
self.vllm_config.parallel_config.enable_expert_parallel = False
self.vllm_config.speculative_config.draft_tensor_parallel_size = 1
self.vllm_config.speculative_config.num_speculative_tokens = 2
self.vllm_config.speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(2)])
self.vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
self.vllm_config.speculative_config.draft_model_config.uses_mrope = False
self.vllm_config.speculative_config.disable_padded_drafter_batch = False
self.vllm_config.speculative_config.parallel_drafting = False
self.vllm_config.speculative_config.draft_parallel_config.tensor_parallel_size = 1
self.vllm_config.speculative_config.target_parallel_config.tensor_parallel_size = 1
self.vllm_config.additional_config = None
init_ascend_config(self.vllm_config)
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
# Set the current vllm config
with set_current_vllm_config(self.vllm_config):
self.proposer = AscendDraftModelProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.proposer.draft_attn_groups = [MagicMock()]
def tearDown(self):
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
# Clear the current vllm config
set_current_vllm_config(None)
@patch('torch.ops._C_ascend.npu_copy_and_expand_eagle_inputs', create=True)
@patch("vllm_ascend.spec_decode.llm_base_proposer.compute_new_slot_mapping")
def test_set_inputs_first_pass(self, mock_slot, mock_expand):
self.assertTrue(self.proposer.needs_extra_input_slots)
target_token_ids = torch.tensor([0,1,2,3,4])
target_positions = torch.tensor([0,1,2,3,4])
next_token_ids = torch.tensor([5])
target_hidden_states = None
token_indices_to_sample = None
num_rejected_tokens_gpu = torch.tensor([0])
batch_size = 1
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=torch.tensor([0, 5], dtype=torch.int32),
query_start_loc_cpu=torch.tensor([0, 5], dtype=torch.int32),
seq_lens=torch.tensor([5], dtype=torch.int32),
seq_lens_cpu=torch.tensor([5], dtype=torch.int32),
num_actual_tokens=5,
max_query_len=5,
max_seq_len=5,
num_reqs=1,
block_table_tensor=torch.zeros([1,320], dtype=torch.int32),
slot_mapping=torch.tensor([128,129,130,131], dtype=torch.int32),
)
common_attn_metadata.batch_size = lambda: batch_size
mock_expand.return_value = (
next_token_ids,
torch.tensor([5]),
torch.tensor([False]),
torch.tensor([False]),
token_indices_to_sample,
None,
)
mock_slot.return_value = torch.tensor([[1]])
_, _, common_attn_metadata, _ = (
self.proposer.set_inputs_first_pass(
target_token_ids,
next_token_ids,
target_positions,
target_hidden_states,
token_indices_to_sample,
common_attn_metadata,
num_rejected_tokens_gpu
)
)
assert common_attn_metadata.seq_lens.to("cpu") == common_attn_metadata.seq_lens_cpu
# fmt: on
class TestEagleProposerPrepareInputs:
"""Test prepare_inputs for AscendEagleProposer.
This test class covers prepare_inputs which handles rejected tokens
and computes token indices for the speculator.
"""
@pytest.fixture(autouse=True)
def setUp_and_tearDown(self):
self.device = torch.device(current_platform.device_type)
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.max_num_tokens = 8192
self.runner.max_num_reqs = 256
self.runner.attn_state = AscendAttentionState.ChunkedPrefill
self.runner.decode_token_per_req = 1
self.runner.actual_seq_lengths_q = []
self.runner.pcp_manager = None
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
yield
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
def _create_base_vllm_config(self):
vllm_config = MagicMock(spec=VllmConfig)
vllm_config.cache_config = MagicMock(spec=CacheConfig)
vllm_config.cache_config.block_size = BLOCK_SIZE
vllm_config.scheduler_config = MagicMock()
vllm_config.scheduler_config.max_num_batched_tokens = 1024
vllm_config.scheduler_config.max_num_seqs = 32
vllm_config.scheduler_config.async_scheduling = False
vllm_config.model_config = MagicMock()
vllm_config.model_config.hf_text_config = MagicMock(spec=[])
vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
vllm_config.model_config.dtype = torch.float16
vllm_config.model_config.max_model_len = 2048
vllm_config.model_config.uses_mrope = False
vllm_config.model_config.uses_xdrope_dim = 0
vllm_config.compilation_config = MagicMock()
vllm_config.parallel_config = MagicMock()
vllm_config.parallel_config.tensor_parallel_size = 1
vllm_config.parallel_config.data_parallel_rank = 0
vllm_config.parallel_config.data_parallel_size = 1
vllm_config.parallel_config.prefill_context_parallel_size = 1
vllm_config.parallel_config.enable_expert_parallel = False
vllm_config.additional_config = {}
return vllm_config
def _create_speculative_config(self, method: str, num_speculative_tokens: int):
speculative_config = MagicMock()
speculative_config.method = method
speculative_config.parallel_drafting = False
speculative_config.num_speculative_tokens = num_speculative_tokens
speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(num_speculative_tokens)])
speculative_config.draft_tensor_parallel_size = 1
speculative_config.disable_padded_drafter_batch = False
speculative_config.draft_model_config = MagicMock()
speculative_config.draft_model_config.get_hidden_size.return_value = 4096
speculative_config.draft_model_config.hf_config.hc_mult = 1
speculative_config.draft_model_config.uses_xdrope_dim = 0
speculative_config.draft_model_config.uses_mrope = False
speculative_config.target_parallel_config = MagicMock()
speculative_config.target_parallel_config.tensor_parallel_size = 1
speculative_config.draft_parallel_config = MagicMock()
speculative_config.draft_parallel_config.tensor_parallel_size = 1
return speculative_config
def _create_proposer(self, method: str, num_speculative_tokens: int, device: torch.device = None, runner=None):
if device is None:
device = torch.device(current_platform.device_type)
vllm_config = self._create_base_vllm_config()
vllm_config.speculative_config = self._create_speculative_config(
method=method,
num_speculative_tokens=num_speculative_tokens,
)
init_ascend_config(vllm_config)
with (
patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer),
patch("vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False),
set_current_vllm_config(vllm_config),
):
proposer = AscendEagleProposer(
vllm_config=vllm_config,
device=device,
runner=runner,
)
proposer.block_size = BLOCK_SIZE
return proposer, vllm_config
def test_prepare_inputs_basic(self):
"""Test prepare_inputs_padded with basic scenario.
Setup:
- 3 requests with query_lens [4, 7, 5]
- num_draft_tokens = [3, 6, 4]
- sampled_token_ids lengths = [2, 3, 2]
- Expected: token_indices should be [0,1, 4,5,6 11,12]
"""
num_speculative_tokens = 6
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
# Setup token_arange_np
proposer.token_arange_np = np.arange(8192)
batch_spec = BatchSpec(
seq_lens=[10, 8, 12],
query_lens=[4, 7, 5],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
)
old_slot_mapping = common_attn_metadata.slot_mapping.clone()
# Define token types
ACCEPT_TOKEN = 0
BONUS_TOKEN = 1
REJECT_TOKEN = -1
sampled_token_ids_with_markers = [
[ACCEPT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
[ACCEPT_TOKEN, ACCEPT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
[ACCEPT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
]
# Filter out rejected tokens
sampled_token_ids = [
[100 * (j + 1) + i for i, token in enumerate(seq) if token != REJECT_TOKEN]
for j, seq in enumerate(sampled_token_ids_with_markers)
]
num_draft_tokens = [3, 6, 4] # one less than query_lens
spec_common_attn_metadata, token_indices = proposer.prepare_inputs(
common_attn_metadata=common_attn_metadata,
sampled_token_ids=sampled_token_ids,
num_draft_tokens=num_draft_tokens,
)
# Actually valid tokens should be 7 (2+3+2)
assert spec_common_attn_metadata.num_actual_tokens == 7
# Accepted token indices
expected_token_indices = torch.tensor(
[0, 1, 4, 5, 6, 11, 12],
dtype=torch.int32,
device=self.device,
)
assert torch.equal(token_indices, expected_token_indices)
# assert attrs computed by prepare_inputs()
attrs_from_prepare_inputs = [
"query_start_loc_cpu",
"query_start_loc",
"seq_lens_cpu",
"_seq_lens_cpu",
"seq_lens",
"seq_lens_cpu_upper_bound",
"num_actual_tokens",
"max_query_len",
"max_seq_len",
"slot_mapping",
]
expected_spec_cad = MagicMock()
# num_tokens_per_req = [2, 3, 2] (after subtracting rejected tokens)
# valid tokens cumulative sum: [0, 2, 5, 7]
expected_spec_cad.query_start_loc_cpu = torch.tensor([0, 2, 5, 7], dtype=torch.int32)
expected_spec_cad.query_start_loc = expected_spec_cad.query_start_loc_cpu.to(self.device, non_blocking=True)
# seq_lens subtract rejected and bonus tokens
expected_spec_cad.seq_lens_cpu = torch.tensor([8, 4, 9], dtype=torch.int32)
expected_spec_cad._seq_lens_cpu = expected_spec_cad.seq_lens_cpu
expected_spec_cad.seq_lens = expected_spec_cad.seq_lens_cpu.to(self.device, non_blocking=True)
expected_spec_cad.seq_lens_cpu_upper_bound = expected_spec_cad.seq_lens_cpu
# actually accepted tokens numble
expected_spec_cad.num_actual_tokens = 7
expected_spec_cad.max_query_len = 7
# default setting is 0
expected_spec_cad.max_seq_len = 0
# compute slot_mapping according to prepare_inputs()
old_slot_mapping[: expected_token_indices.shape[0]].copy_(old_slot_mapping[expected_token_indices])
old_slot_mapping[expected_token_indices.shape[0] :].fill_(-1)
expected_spec_cad.slot_mapping = old_slot_mapping
for attr in attrs_from_prepare_inputs:
assert_attr_equal(attr, expected_spec_cad, spec_common_attn_metadata)
# assert attrs inherited from common_attn_metadata
attrs_from_cad: list[str | tuple[str, Any, Any]] = [
"num_computed_tokens_cpu",
"_num_computed_tokens_cpu",
"num_reqs",
"num_input_tokens",
"block_table_tensor",
"slot_mapping",
("positions", expected_token_indices, None),
]
for metadata_attr in attrs_from_cad:
assert_attr_equal(metadata_attr, common_attn_metadata, spec_common_attn_metadata)
# assert attrs inherited from runner
attrs_from_runner = [
"actual_seq_lengths_q",
"attn_state",
"decode_token_per_req",
]
for attr in attrs_from_runner:
assert_attr_equal(attr, self.runner, spec_common_attn_metadata)
def test_prepare_inputs_all_rejected(self):
"""Test prepare_inputs when all the tokens are rejected.
Setup:
- 3 requests with query_lens [4, 3, 5]
- num_draft_tokens = [3, 2, 4]
- sampled_token_ids lengths = [1, 1, 1]
- num_rejected = [3, 2, 4] # bonus tokens included
"""
num_speculative_tokens = 4
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
proposer.token_arange_np = np.arange(8192)
batch_spec = BatchSpec(
seq_lens=[10, 8, 12],
query_lens=[4, 3, 5],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
)
old_slot_mapping = common_attn_metadata.slot_mapping.clone()
# Define token types
BONUS_TOKEN = 1
REJECT_TOKEN = -1
# still sample one bonus_token though all rejected
sampled_token_ids_with_markers = [
[REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
[REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
[REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, REJECT_TOKEN, BONUS_TOKEN],
]
# Filter out rejected tokens
sampled_token_ids = [
[100 + i * 10 for i, token in enumerate(seq) if token != REJECT_TOKEN]
for seq in sampled_token_ids_with_markers
]
num_draft_tokens = [3, 2, 4]
spec_common_attn_metadata, token_indices = proposer.prepare_inputs(
common_attn_metadata=common_attn_metadata,
sampled_token_ids=sampled_token_ids,
num_draft_tokens=num_draft_tokens,
)
# no tokens accepted, just sample bonus tokens
expected_token_indices = torch.tensor(
[0, 4, 7],
dtype=torch.int32,
device=self.device,
)
assert torch.equal(token_indices, expected_token_indices)
# assert attrs computed by prepare_inputs()
attrs_from_prepare_inputs = [
"query_start_loc_cpu",
"query_start_loc",
"seq_lens_cpu",
"_seq_lens_cpu",
"seq_lens",
"seq_lens_cpu_upper_bound",
"num_actual_tokens",
"max_query_len",
"max_seq_len",
"slot_mapping",
]
expected_spec_cad = MagicMock()
# num_tokens_per_req = [1, 1, 1] (query_lens subtracting rejected tokens)
# valid tokens cumulative sum: [0, 1, 2, 3]
expected_spec_cad.query_start_loc_cpu = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
expected_spec_cad.query_start_loc = expected_spec_cad.query_start_loc_cpu.to(self.device, non_blocking=True)
# seq_lens subtract rejected and bonus tokens
expected_spec_cad.seq_lens_cpu = torch.tensor([7, 6, 8], dtype=torch.int32)
expected_spec_cad._seq_lens_cpu = expected_spec_cad.seq_lens_cpu
expected_spec_cad.seq_lens = expected_spec_cad.seq_lens_cpu.to(self.device, non_blocking=True)
expected_spec_cad.seq_lens_cpu_upper_bound = expected_spec_cad.seq_lens_cpu
# actually accepted tokens numble
expected_spec_cad.num_actual_tokens = 3
expected_spec_cad.max_query_len = 5
# default setting is 0
expected_spec_cad.max_seq_len = 0
# compute slot_mapping according to prepare_inputs()
old_slot_mapping[: expected_token_indices.shape[0]].copy_(old_slot_mapping[expected_token_indices])
old_slot_mapping[expected_token_indices.shape[0] :].fill_(-1)
expected_spec_cad.slot_mapping = old_slot_mapping
for attr in attrs_from_prepare_inputs:
assert_attr_equal(attr, expected_spec_cad, spec_common_attn_metadata)
# assert attrs inherited from common_attn_metadata
attrs_from_cad: list[str | tuple[str, Any, Any]] = [
"num_computed_tokens_cpu",
"_num_computed_tokens_cpu",
"num_reqs",
"num_input_tokens",
"block_table_tensor",
"slot_mapping",
("positions", expected_token_indices, None),
]
for metadata_attr in attrs_from_cad:
assert_attr_equal(metadata_attr, common_attn_metadata, spec_common_attn_metadata)
# assert attrs inherited from runner
attrs_from_runner = [
"actual_seq_lengths_q",
"attn_state",
"decode_token_per_req",
]
for attr in attrs_from_runner:
assert_attr_equal(attr, self.runner, spec_common_attn_metadata)
class TestEagleProposerPrepareInputsPadded:
"""Test prepare_inputs_padded for AscendEagleProposer.
This test class covers prepare_inputs_padded which handles padded inputs
for speculative decoding without considering rejected tokens.
"""
@pytest.fixture(autouse=True)
def setUp_and_tearDown(self):
self.device = torch.device(current_platform.device_type)
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.max_num_tokens = 8192
self.runner.max_num_reqs = 256
self.runner.attn_state = AscendAttentionState.ChunkedPrefill
self.runner.decode_token_per_req = 1
self.runner.actual_seq_lengths_q = []
self.runner.pcp_manager = None
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
yield
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
def _create_base_vllm_config(self):
vllm_config = MagicMock(spec=VllmConfig)
vllm_config.cache_config = MagicMock(spec=CacheConfig)
vllm_config.cache_config.block_size = BLOCK_SIZE
vllm_config.scheduler_config = MagicMock()
vllm_config.scheduler_config.max_num_batched_tokens = 1024
vllm_config.scheduler_config.max_num_seqs = 32
vllm_config.scheduler_config.async_scheduling = False
vllm_config.model_config = MagicMock()
vllm_config.model_config.hf_text_config = MagicMock(spec=[])
vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
vllm_config.model_config.dtype = torch.float16
vllm_config.model_config.max_model_len = 2048
vllm_config.model_config.uses_mrope = False
vllm_config.model_config.uses_xdrope_dim = 0
vllm_config.compilation_config = MagicMock()
vllm_config.parallel_config = MagicMock()
vllm_config.parallel_config.tensor_parallel_size = 1
vllm_config.parallel_config.data_parallel_rank = 0
vllm_config.parallel_config.data_parallel_size = 1
vllm_config.parallel_config.prefill_context_parallel_size = 1
vllm_config.parallel_config.enable_expert_parallel = False
vllm_config.additional_config = {}
return vllm_config
def _create_speculative_config(self, method: str, num_speculative_tokens: int):
speculative_config = MagicMock()
speculative_config.method = method
speculative_config.parallel_drafting = False
speculative_config.num_speculative_tokens = num_speculative_tokens
speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(num_speculative_tokens)])
speculative_config.draft_tensor_parallel_size = 1
speculative_config.disable_padded_drafter_batch = False
speculative_config.draft_model_config = MagicMock()
speculative_config.draft_model_config.get_hidden_size.return_value = 4096
speculative_config.draft_model_config.hf_config.hc_mult = 1
speculative_config.draft_model_config.uses_xdrope_dim = 0
speculative_config.draft_model_config.uses_mrope = False
speculative_config.target_parallel_config = MagicMock()
speculative_config.target_parallel_config.tensor_parallel_size = 1
speculative_config.draft_parallel_config = MagicMock()
speculative_config.draft_parallel_config.tensor_parallel_size = 1
return speculative_config
def _create_proposer(self, method: str, num_speculative_tokens: int, device: torch.device = None, runner=None):
if device is None:
device = torch.device(current_platform.device_type)
vllm_config = self._create_base_vllm_config()
vllm_config.speculative_config = self._create_speculative_config(
method=method,
num_speculative_tokens=num_speculative_tokens,
)
init_ascend_config(vllm_config)
with (
patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer),
patch("vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False),
set_current_vllm_config(vllm_config),
):
proposer = AscendEagleProposer(
vllm_config=vllm_config,
device=device,
runner=runner,
)
proposer.block_size = BLOCK_SIZE
return proposer, vllm_config
@pytest.mark.parametrize(
"has_triton,num_aicore,num_vectorcore",
[
(True, 24, 48),
(False, -1, -1),
],
)
def test_prepare_inputs_padded_basic(self, has_triton, num_aicore, num_vectorcore):
"""Test prepare_inputs_padded with basic scenario.
Setup:
- 3 requests with query_lens [5, 5, 5]
- num_draft_tokens = [4, 4, 4]
- cu_num_draft_tokens: [4, 8, 12]
- num_rejected_tokens: [3, 2, 1]
- valid_sampled_tokens_count: [2, 3, 4]
"""
num_speculative_tokens = 4
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
proposer.arange = torch.arange(8192, dtype=torch.int32, device=self.device)
batch_spec = BatchSpec(
seq_lens=[10, 8, 12],
query_lens=[5, 5, 5],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
)
# Mock SpecDecodeMetadata (inclusive cumsum, first element is first value not zero)
# Inclusive cumsum: [4, 8, 12] means all reqs have 4 draft tokens
spec_decode_metadata = MagicMock()
spec_decode_metadata.cu_num_draft_tokens = torch.tensor([4, 8, 12], dtype=torch.int32, device=self.device)
valid_sampled_tokens_count = torch.tensor([2, 3, 4], dtype=torch.int32, device=self.device)
with (
patch(
"vllm_ascend.spec_decode.llm_base_proposer.HAS_TRITON",
has_triton,
),
patch.multiple(
"vllm_ascend.ops.triton.triton_utils",
_NUM_AICORE=num_aicore,
_NUM_VECTORCORE=num_vectorcore,
),
):
spec_common_attn_metadata, token_indices, token_indices_to_sample, num_rejected_tokens_gpu = (
proposer.prepare_inputs_padded(
common_attn_metadata=common_attn_metadata,
spec_decode_metadata=spec_decode_metadata,
valid_sampled_tokens_count=valid_sampled_tokens_count,
)
)
# Total tokens should be 15
assert token_indices.shape[0] == 15
expected_token_indices = torch.arange(15, dtype=torch.int32, device=self.device)
assert torch.equal(token_indices, expected_token_indices)
# num_rejected_tokens_gpu
# req0: 4+1-2 = 3
# req1: 4+1-3 = 2
# req2: 4+1-4 = 1
expected_num_rejected = torch.tensor([3, 2, 1], dtype=torch.int32, device=self.device)
assert torch.equal(num_rejected_tokens_gpu, expected_num_rejected)
# token_indices_to_sample should be at the end of valid tokens of each request
# firstly subtract one to get the end of each request, and then subtract the rejected ones
# req0: query_start_loc[1]-1-rejected = 5-1-3 = 1
# req1: 10-1-2 = 7
# req2: 15-1-1 = 13
expected_token_indices_to_sample = torch.tensor([1, 7, 13], dtype=torch.int32, device=self.device)
assert torch.equal(token_indices_to_sample, expected_token_indices_to_sample)
# assert attrs inherited from common_attn_metadata
attrs_from_cad = [
"query_start_loc",
"seq_lens_cpu",
"_seq_lens_cpu",
"seq_lens_cpu_upper_bound",
"num_reqs",
"num_input_tokens",
"block_table_tensor",
"slot_mapping",
"positions",
"num_computed_tokens_cpu",
"_num_computed_tokens_cpu",
"seq_lens",
]
for attr in attrs_from_cad:
assert_attr_equal(attr, common_attn_metadata, spec_common_attn_metadata)
attrs_from_prepare_inputs = [
"num_actual_tokens",
"max_query_len",
"max_seq_len",
]
expected_spec_cad = MagicMock()
expected_spec_cad.num_actual_tokens = 15
expected_spec_cad.max_query_len = 5
expected_spec_cad.max_seq_len = 0
for attr in attrs_from_prepare_inputs:
assert_attr_equal(attr, expected_spec_cad, spec_common_attn_metadata)
# assert attrs inherited from runner
attrs_from_runner = [
"actual_seq_lengths_q",
"attn_state",
"decode_token_per_req",
]
for attr in attrs_from_runner:
assert_attr_equal(attr, self.runner, spec_common_attn_metadata)
@pytest.mark.parametrize(
"has_triton,num_aicore,num_vectorcore",
[
(True, 24, 48),
(False, -1, -1),
],
)
def test_prepare_inputs_padded_all_rejected(self, has_triton, num_aicore, num_vectorcore):
"""Test prepare_inputs_padded when all draft tokens are rejected.
Setup:
- 2 requests with query_lens [4, 3, 5]
- num_draft_tokens = [3, 2, 4]
- cu_num_draft_tokens: [3, 5, 9]
- num_rejected_tokens: [3, 2, 4]
- valid_sampled_tokens_count: [1, 1, 1] (only bonus token)
"""
num_speculative_tokens = 4
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
proposer.arange = torch.arange(8192, dtype=torch.int32, device=self.device)
batch_spec = BatchSpec(
seq_lens=[10, 8, 12],
query_lens=[4, 3, 5],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
)
# Mock SpecDecodeMetadata (inclusive cumsum, first element is first value not zero)
# Inclusive cumsum: [3, 5, 9] means req0 has 3, req1 has 2 draft tokens, req2 has 4 draft tokens
spec_decode_metadata = MagicMock()
spec_decode_metadata.cu_num_draft_tokens = torch.tensor([3, 5, 9], dtype=torch.int32, device=self.device)
valid_sampled_tokens_count = torch.tensor([1, 1, 1], dtype=torch.int32, device=self.device)
with (
patch(
"vllm_ascend.spec_decode.llm_base_proposer.HAS_TRITON",
has_triton,
),
patch.multiple(
"vllm_ascend.ops.triton.triton_utils",
_NUM_AICORE=num_aicore,
_NUM_VECTORCORE=num_vectorcore,
),
):
spec_common_attn_metadata, token_indices, token_indices_to_sample, num_rejected_tokens_gpu = (
proposer.prepare_inputs_padded(
common_attn_metadata=common_attn_metadata,
spec_decode_metadata=spec_decode_metadata,
valid_sampled_tokens_count=valid_sampled_tokens_count,
)
)
# Total tokens: 4 + 3 + 5 = 12
assert token_indices.shape[0] == 12
expected_token_indices = torch.arange(12, dtype=torch.int32, device=self.device)
assert torch.equal(token_indices, expected_token_indices)
# num_rejected_tokens_gpu
# req0: 3+1-1 = 3
# req1: 2+1-1 = 2
# req2: 4+1-1 = 4
expected_num_rejected = torch.tensor([3, 2, 4], dtype=torch.int32, device=self.device)
assert torch.equal(num_rejected_tokens_gpu, expected_num_rejected)
# token_indices_to_sample should be at the end of valid tokens of each request
# firstly subtract one to get the end of each request, and then subtract the rejected ones
# req0: query_start_loc[1]-1-rejected = 4-1-3 = 0
# req1: 7-1-2 = 4
# req2: 12-1-4 = 7
expected_token_indices_to_sample = torch.tensor([0, 4, 7], dtype=torch.int32, device=self.device)
assert torch.equal(token_indices_to_sample, expected_token_indices_to_sample)
# assert attrs inherited from common_attn_metadata
attrs_from_cad = [
"query_start_loc",
"seq_lens_cpu",
"_seq_lens_cpu",
"seq_lens_cpu_upper_bound",
"num_reqs",
"num_input_tokens",
"block_table_tensor",
"slot_mapping",
"positions",
"num_computed_tokens_cpu",
"_num_computed_tokens_cpu",
"seq_lens",
]
for attr in attrs_from_cad:
assert_attr_equal(attr, common_attn_metadata, spec_common_attn_metadata)
# assert attrs computed by prepare_inputs()
attrs_from_prepare_inputs = [
"num_actual_tokens",
"max_query_len",
"max_seq_len",
]
expected_spec_cad = MagicMock()
expected_spec_cad.num_actual_tokens = 12
expected_spec_cad.max_query_len = 5
expected_spec_cad.max_seq_len = 0
for attr in attrs_from_prepare_inputs:
assert_attr_equal(attr, expected_spec_cad, spec_common_attn_metadata)
# assert attrs inherited from runner
attrs_from_runner = [
"actual_seq_lengths_q",
"attn_state",
"decode_token_per_req",
]
for attr in attrs_from_runner:
assert_attr_equal(attr, self.runner, spec_common_attn_metadata)
class TestEagleProposerSetInputsFirstPass:
"""Test set_inputs_first_pass for AscendEagleProposer.
This test class covers all branches of set_inputs_first_pass:
Branch coverage:
- Branch 1 (needs_extra_input_slots=False): Default EAGLE pathway
- Branch 1.1: multiple requests
- Branch 1.2: pcp_size > 1 (PCP split logic) - vllm-ascend specific
- Branch 2 (needs_extra_input_slots=True): Draft model / Parallel drafting
- Branch 2.1: shift_input_ids=False (draft_model)
- Branch 2.2: shift_input_ids=True (parallel_drafting)
"""
@pytest.fixture(autouse=True)
def setUp_and_tearDown(self):
self.device = torch.device(current_platform.device_type)
self.runner = MagicMock()
self.runner.pin_memory = False
self.runner.pcp_size = 1
self.runner.dcp_size = 1
self.runner.max_num_tokens = 8192
self.runner.max_num_reqs = 256
self.runner.pcp_manager = None
self.mock_cpugpubuffer = patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer)
self.mock_cpugpubuffer.start()
self.mock_supports_multimodal_inputs = patch(
"vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False
)
self.mock_supports_multimodal_inputs.start()
yield
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
def _create_base_vllm_config(self):
"""Create base vllm_config with common settings shared across all tests."""
vllm_config = MagicMock(spec=VllmConfig)
vllm_config.cache_config = MagicMock(spec=CacheConfig)
vllm_config.cache_config.block_size = BLOCK_SIZE
vllm_config.scheduler_config = MagicMock()
vllm_config.scheduler_config.max_num_batched_tokens = 1024
vllm_config.scheduler_config.max_num_seqs = 32
vllm_config.scheduler_config.async_scheduling = False
vllm_config.model_config = MagicMock()
vllm_config.model_config.hf_text_config = MagicMock(spec=[])
vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
vllm_config.model_config.dtype = torch.float16
vllm_config.model_config.max_model_len = 2048
vllm_config.model_config.uses_mrope = False
vllm_config.model_config.uses_xdrope_dim = 0
vllm_config.compilation_config = MagicMock()
vllm_config.parallel_config = MagicMock()
vllm_config.parallel_config.tensor_parallel_size = 1
vllm_config.parallel_config.data_parallel_rank = 0
vllm_config.parallel_config.data_parallel_size = 1
vllm_config.parallel_config.prefill_context_parallel_size = 1
vllm_config.parallel_config.enable_expert_parallel = False
vllm_config.additional_config = {}
return vllm_config
def _create_speculative_config(
self,
method: str,
num_speculative_tokens: int,
parallel_drafting: bool = False,
):
"""Create speculative_config for specific method."""
speculative_config = MagicMock()
speculative_config.method = method
speculative_config.parallel_drafting = parallel_drafting
speculative_config.num_speculative_tokens = num_speculative_tokens
speculative_config.speculative_token_tree = str([(i + 1) * (0,) for i in range(num_speculative_tokens)])
speculative_config.draft_tensor_parallel_size = 1
speculative_config.disable_padded_drafter_batch = False
speculative_config.draft_model_config = MagicMock()
speculative_config.draft_model_config.get_hidden_size.return_value = 4096
speculative_config.draft_model_config.hf_config.hc_mult = 1
speculative_config.draft_model_config.uses_xdrope_dim = 0
speculative_config.draft_model_config.uses_mrope = False
speculative_config.target_parallel_config = MagicMock()
speculative_config.target_parallel_config.tensor_parallel_size = 1
speculative_config.draft_parallel_config = MagicMock()
speculative_config.draft_parallel_config.tensor_parallel_size = 1
return speculative_config
def _create_proposer(
self,
method: str,
num_speculative_tokens: int,
parallel_drafting: bool = False,
device: torch.device = None,
runner=None,
):
"""Create a proposer instance for testing."""
if device is None:
device = torch.device(current_platform.device_type)
vllm_config = self._create_base_vllm_config()
vllm_config.speculative_config = self._create_speculative_config(
method=method,
num_speculative_tokens=num_speculative_tokens,
parallel_drafting=parallel_drafting,
)
init_ascend_config(vllm_config)
with (
patch(_CPU_GPU_BUFFER_TARGET, MockCpuGpuBuffer),
patch("vllm.multimodal.registry.MultiModalRegistry.supports_multimodal_inputs", return_value=False),
set_current_vllm_config(vllm_config),
):
if method == "eagle":
proposer = AscendEagleProposer(
vllm_config=vllm_config,
device=device,
runner=runner,
)
elif method == "draft_model":
proposer = AscendDraftModelProposer(
vllm_config=vllm_config,
device=device,
runner=runner,
)
proposer.block_size = BLOCK_SIZE
return proposer, vllm_config
def test_set_inputs_first_pass_default_eagle(self):
"""
Test for set_inputs_first_pass without extra input slots (default EAGLE).
This tests the path where needs_extra_input_slots=False, which is the
default EAGLE pathway. In this case:
- Input IDs are rotated (shifted by one)
- The next_token_ids are inserted at the last position of each request
- Positions are copied as-is
- Hidden states are copied as-is
- The CommonAttentionMetadata is returned unchanged
Setup:
- 3 requests with query_lens [3, 2, 4]
- Tokens: [a1, a2, a3, b1, b2, c1, c2, c3, c4]
- After rotation: [a2, a3, -, b2, -, c2, c3, c4, -]
- After inserting next_tokens [100, 200, 300]:
[a2, a3, 100, b2, 200, c2, c3, c4, 300]
"""
num_speculative_tokens = 3
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
batch_spec = BatchSpec(
seq_lens=[10, 8, 12],
query_lens=[3, 2, 4],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
)
target_token_ids = torch.tensor([10, 11, 12, 20, 21, 30, 31, 32, 33], dtype=torch.int32, device=self.device)
target_positions = torch.tensor([7, 8, 9, 6, 7, 8, 9, 10, 11], dtype=torch.int64, device=self.device)
target_hidden_states = torch.randn(9, proposer.hidden_size, dtype=proposer.dtype, device=self.device)
next_token_ids = torch.tensor([100, 200, 300], dtype=torch.int32, device=self.device)
out_num_tokens, out_token_indices, out_cad, long_seq_args = proposer.set_inputs_first_pass(
target_token_ids=target_token_ids,
next_token_ids=next_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
token_indices_to_sample=None,
cad=common_attn_metadata,
num_rejected_tokens_gpu=None,
)
# assert function computed outputs
assert out_num_tokens == 9
expected_token_indices = torch.tensor([2, 4, 8], dtype=torch.int32, device=self.device)
assert torch.equal(out_token_indices, expected_token_indices)
assert out_cad is common_attn_metadata # returned as-is
assert long_seq_args is None
# assert proposer internal state
expected_proposer = MagicMock()
expected_proposer.input_ids = torch.tensor(
[11, 12, 100, 21, 200, 31, 32, 33, 300], dtype=torch.int32, device=self.device
)
expected_proposer.positions = target_positions
expected_proposer.hidden_states = target_hidden_states
attrs_from_proposer: list[str | tuple[str, Any, Any]] = [
("input_ids", None, slice(None, out_num_tokens)),
("positions", None, slice(None, out_num_tokens)),
("hidden_states", None, slice(None, out_num_tokens)),
]
for attr in attrs_from_proposer:
assert_attr_equal(attr, expected_proposer, proposer)
def test_set_inputs_first_pass_pcp_dcp_mixed(self):
"""
Test Default pcp_dcp_mixed scenario
Just for coverage no reference value
Maybe rewrite considering the refactor of pcp-dcp
"""
num_speculative_tokens = 3
block_size = BLOCK_SIZE
req_ids = ["req-0", "req-1", "req-2", "req-3"]
req_scheduled_tokens = {"req-0": 3, "req-1": 2, "req-2": 4, "req-3": 3}
query_lens = [3, 2, 4, 3]
self.runner.query_lens = torch.tensor(query_lens, dtype=torch.int32, device=self.device)
self.runner.input_batch = MagicMock()
self.runner.input_batch.req_ids = req_ids
# maybe not reasonable just to run test
self.runner.logits_indices = torch.arange(12, dtype=torch.int32, device=self.device)
pcp_manager = MagicMock()
self.runner.pcp_manager = pcp_manager
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
proposer.pcp_size = 2
proposer.dcp_size = 2
proposer.pcp_rank = 0
proposer.needs_extra_input_slots = False
num_decode_reqs = 2
num_prefill_reqs = 2
batch_spec = BatchSpec(
seq_lens=[10, 8, 12, 6],
query_lens=query_lens,
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
arange_block_indices=True,
)
target_token_ids = torch.tensor(
[10, 11, 12, 20, 21, 30, 31, 32, 33, 40, 41, 42], dtype=torch.int32, device=self.device
)
target_positions = torch.tensor([7, 8, 9, 6, 7, 8, 9, 10, 11, 5, 6, 7], dtype=torch.int64, device=self.device)
next_token_ids = torch.tensor([100, 200, 300, 400], dtype=torch.int32, device=self.device)
target_hidden_states = torch.randn(18, proposer.hidden_size, dtype=proposer.dtype, device=self.device)
long_seq_metadata = MagicMock()
expected_token_indices = torch.tensor([2, 4, 10, 11], dtype=torch.int32, device=self.device)
expected_query_lens_d = torch.tensor([3, 2], dtype=torch.int32, device=self.device)
expected_ori_token_indices_to_sample = torch.tensor([2, 4, 8, 11], dtype=torch.int32, device=self.device)
expected_input_ids = torch.tensor(
[11, 12, 100, 21, 200, 31, 300, 41, 0],
dtype=torch.int32,
device=self.device,
)
expected_positions = target_positions[:9]
expected_hidden_indices = torch.tensor([0, 1, 2, 6, 7, 10, 13, 14, 17], dtype=torch.long, device=self.device)
expected_hidden_states = target_hidden_states[expected_hidden_indices]
def prepare_first_pass_inputs(**kwargs):
assert torch.equal(
kwargs["input_ids"],
torch.tensor(
[11, 12, 100, 21, 200, 31, 32, 33, 300, 41, 42, 400],
dtype=torch.int32,
device=self.device,
),
)
assert kwargs["common_attn_metadata"] is common_attn_metadata
assert kwargs["long_seq_metadata"] is long_seq_metadata
assert kwargs["req_scheduled_tokens"] == req_scheduled_tokens
assert kwargs["req_ids"] == req_ids
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
common_attn_metadata.num_actual_tokens = 9
common_attn_metadata.seq_lens = torch.tensor([10, 8, 2, 2], dtype=torch.int32, device=self.device)
common_attn_metadata.query_start_loc = torch.tensor([0, 3, 5, 7, 9], dtype=torch.int32, device=self.device)
common_attn_metadata.max_query_len = 4
return PCPSpecDecodeFirstPassInputs(
num_tokens=9,
input_ids=expected_input_ids,
target_positions=expected_positions,
target_hidden_states=expected_hidden_states,
token_indices_to_sample=expected_token_indices,
long_seq_args=(expected_query_lens_d, expected_ori_token_indices_to_sample),
)
pcp_manager.prepare_spec_decode_first_pass_inputs.side_effect = prepare_first_pass_inputs
out_num_tokens, out_token_indices, out_cad, (query_lens_d, ori_token_indices_to_sample) = (
proposer.set_inputs_first_pass(
target_token_ids=target_token_ids,
next_token_ids=next_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
token_indices_to_sample=None,
cad=common_attn_metadata,
num_rejected_tokens_gpu=None,
req_scheduled_tokens=req_scheduled_tokens,
long_seq_metadata=long_seq_metadata,
num_prefill_reqs=num_prefill_reqs,
num_decode_reqs=num_decode_reqs,
)
)
# assert function computed outputs
assert out_num_tokens == 9
assert torch.equal(out_token_indices, expected_token_indices)
pcp_manager.prepare_spec_decode_first_pass_inputs.assert_called_once()
# assert query_lens_d and ori_token_indices_to_sample
assert torch.equal(query_lens_d, expected_query_lens_d)
assert torch.equal(ori_token_indices_to_sample, expected_ori_token_indices_to_sample)
# assert proposer internal state
expected_proposer = MagicMock()
expected_proposer.input_ids = expected_input_ids
expected_proposer.positions = expected_positions
expected_proposer.hidden_states = expected_hidden_states
attrs_from_proposer: list[str | tuple[str, Any, Any]] = [
("input_ids", None, slice(None, out_num_tokens)),
("positions", None, slice(None, out_num_tokens)),
("hidden_states", None, slice(None, out_num_tokens)),
]
for attr in attrs_from_proposer:
assert_attr_equal(attr, expected_proposer, proposer)
# assert metadata attributes modified by PCP logic
expected_metadata = MagicMock()
expected_metadata.num_actual_tokens = 9
expected_metadata.seq_lens = torch.tensor([10, 8, 2, 2], dtype=torch.int32, device=self.device)
expected_metadata.query_start_loc = torch.tensor([0, 3, 5, 7, 9], dtype=torch.int32, device=self.device)
expected_metadata.max_query_len = 4
attrs_from_metadata = [
"num_actual_tokens",
"seq_lens",
"query_start_loc",
"max_query_len",
]
for attr in attrs_from_metadata:
assert_attr_equal(attr, expected_metadata, out_cad)
assert out_cad.prefill_context_parallel_metadata == long_seq_metadata
def test_set_inputs_first_pass_parallel_drafting(self):
"""
Test for set_inputs_first_pass with parallel drafting (extra input slots,
with shift).
This tests the path where needs_extra_input_slots=True and
shift_input_ids=True (parallel drafting case). In this case:
- Input IDs ARE shifted (like default EAGLE)
- Each request gets extra_slots_per_request (3) new slots
- Parallel drafting tokens are inserted and marked as masked
- Hidden states are mapped correctly
Setup:
- 2 requests with query_lens [4, 4] (1 bonus + 3 spec tokens each)
- Request 0: tokens [10, 11, 12, 13] at positions [5, 6, 7, 8]
- Only tokens [10, 11, 12] are "valid", token 13 is rejected
- Request 1: tokens [20, 21, 22, 23] at positions [10, 11, 12, 13], all valid.
- next_token_ids: [100, 200] (bonus tokens)
With shift_input_ids=True, extra_slots_per_request=3:
Expected output layout:
Request 0 (6 output slots = 4 - 1 + 3):
- idx 0-2: shifted tokens [11, 12, 100]
- idx 3-4: parallel_drafting_tokens, is_masked=True
- idx 5: padding_token, is_rejected=True
Request 1 (6 output slots = 4 - 1 + 3):
- idx 6-8: shifted tokens [21, 22, 23]
- idx 9: bonus token 200
- idx 10-11: parallel_drafting_tokens, is_masked=True
"""
num_speculative_tokens = 3
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="eagle",
num_speculative_tokens=num_speculative_tokens,
parallel_drafting=True,
device=self.device,
runner=self.runner,
)
proposer.parallel_drafting_token_id = -2
assert proposer.parallel_drafting_hidden_state_tensor is not None
proposer.parallel_drafting_hidden_state_tensor.zero_()
parallel_drafting_hs = proposer.parallel_drafting_hidden_state_tensor
mock_kv_cache_spec = MagicMock()
mock_kv_cache_spec.block_size = block_size
mock_attn_group = MagicMock()
mock_attn_group.kv_cache_spec = mock_kv_cache_spec
proposer.draft_attn_groups = [mock_attn_group]
batch_spec = BatchSpec(
seq_lens=[9, 14],
query_lens=[4, 4],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
arange_block_indices=True,
)
target_token_ids = torch.tensor([10, 11, 12, 13, 20, 21, 22, 23], dtype=torch.int32, device=self.device)
target_positions = torch.tensor([5, 6, 7, 8, 10, 11, 12, 13], dtype=torch.int64, device=self.device)
target_hidden_states = torch.randn(8, proposer.hidden_size, dtype=proposer.dtype, device=self.device).view(
8, proposer.hidden_size
)
next_token_ids = torch.tensor([100, 200], dtype=torch.int32, device=self.device)
num_rejected_tokens_gpu = torch.tensor([1, 0], dtype=torch.int32, device=self.device)
out_num_tokens, out_token_indices, out_cad, long_seq_args = proposer.set_inputs_first_pass(
target_token_ids=target_token_ids,
next_token_ids=next_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
token_indices_to_sample=None,
cad=common_attn_metadata,
num_rejected_tokens_gpu=num_rejected_tokens_gpu,
)
# assert function computed outputs
assert out_num_tokens == 12
expected_out_token_indices = torch.tensor([2, 3, 4, 9, 10, 11], dtype=torch.int32, device=self.device)
assert torch.equal(out_token_indices, expected_out_token_indices)
assert long_seq_args is None
# assert attrs from proposer
attrs_from_proposer: list[tuple[str, Any, Any]] = [
("input_ids", None, slice(None, out_num_tokens)),
("positions", None, slice(None, out_num_tokens)),
("is_rejected_token_mask", None, slice(None, out_num_tokens)),
("is_masked_token_mask", None, slice(None, out_num_tokens)),
("hidden_states", None, slice(None, out_num_tokens)),
]
expected_proposer = MagicMock()
expected_proposer.input_ids = torch.tensor(
[11, 12, 100, -2, -2, 0, 21, 22, 23, 200, -2, -2], dtype=torch.int32, device=self.device
)
expected_proposer.positions = torch.tensor(
[5, 6, 7, 8, 9, 0, 10, 11, 12, 13, 14, 15], dtype=torch.int64, device=self.device
)
expected_proposer.is_rejected_token_mask = torch.tensor(
[0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0], device=self.device, dtype=bool
)
expected_proposer.is_masked_token_mask = torch.tensor(
[0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 1, 1], device=self.device, dtype=bool
)
target_hidden_states_padded = torch.cat(
[
target_hidden_states[:4],
torch.zeros(2, proposer.hidden_size, device=self.device, dtype=proposer.dtype),
target_hidden_states[4:8],
torch.zeros(2, proposer.hidden_size, device=self.device, dtype=proposer.dtype),
],
)
expected_proposer.hidden_states = torch.where(
expected_proposer.is_masked_token_mask.unsqueeze(1), parallel_drafting_hs, target_hidden_states_padded
)
for attr in attrs_from_proposer:
assert_attr_equal(attr, expected_proposer, proposer)
# assert attrs from cad
attrs_from_cad: list[str | tuple[str, Any, Any]] = [
"query_start_loc_cpu",
"query_start_loc",
"seq_lens",
"num_actual_tokens",
"max_query_len",
"max_seq_len",
"slot_mapping",
"seq_lens_cpu",
"_seq_lens_cpu",
]
expected_cad = MagicMock()
expected_cad.query_start_loc_cpu = torch.tensor([0, 6, 12], dtype=torch.int32)
expected_cad.query_start_loc = expected_cad.query_start_loc_cpu.to(self.device, non_blocking=True)
expected_cad.seq_lens = torch.tensor([11, 16], device=self.device, dtype=torch.int32)
expected_cad.num_actual_tokens = 12
expected_cad.max_query_len = 6
expected_cad.max_seq_len = 16
expected_cad.slot_mapping = torch.tensor(
[5, 6, 7, 8, 9, -1, 26, 27, 28, 29, 30, 31], device=self.device, dtype=torch.int64
)
expected_cad.seq_lens_cpu = torch.tensor([11, 16], dtype=torch.int32)
expected_cad._seq_lens_cpu = torch.tensor([11, 16], dtype=torch.int32)
for attrition in attrs_from_cad:
assert_attr_equal(attrition, expected_cad, out_cad)
def test_set_inputs_first_pass_draft_model(self):
"""
Test for set_inputs_first_pass with a draft model (extra input slots,
no shift).
This tests the path where needs_extra_input_slots=True and
shift_input_ids=False (draft model case). In this case:
- Input IDs are NOT shifted
- Each request gets extra_slots_per_request (1) new slots
- The kernel handles copying tokens and inserting bonus/padding tokens
- A new CommonAttentionMetadata is returned with updated query_start_loc
Setup:
- 2 requests
- Request 0: tokens [10, 11, 12] at positions [0, 1, 2]
- Only tokens [10, 11] are "valid" (query_end_loc=1),
token 12 is a rejected token from previous speculation
- Request 1: tokens [20, 21] at positions [0, 1], both valid.
- Note: this is less than num_speculative_tokens (2) to ensure
we handle variable lengths correctly.
- next_token_ids: [100, 200] (bonus tokens)
With extra_slots_per_request=1 and shift=False:
Expected output layout:
Request 0 (indices 0-3):
- idx 0: token 10, pos 0
- idx 1: token 11, pos 1
- idx 2: token 100, pos 2 (bonus token)
- idx 3: padding_token_id, is_rejected=True
Request 1 (indices 4-6):
- idx 4: token 20, pos 0
- idx 5: token 21, pos 1
- idx 6: token 200, pos 2 (bonus token)
"""
num_speculative_tokens = 2
block_size = BLOCK_SIZE
proposer, vllm_config = self._create_proposer(
method="draft_model",
num_speculative_tokens=num_speculative_tokens,
device=self.device,
runner=self.runner,
)
mock_kv_cache_spec = MagicMock()
mock_kv_cache_spec.block_size = block_size
mock_attn_group = MagicMock()
mock_attn_group.kv_cache_spec = mock_kv_cache_spec
proposer.draft_attn_groups = [mock_attn_group]
batch_spec = BatchSpec(
seq_lens=[3, 2],
query_lens=[3, 2],
)
common_attn_metadata = create_common_attn_metadata(
batch_spec,
block_size=block_size,
device=self.device,
arange_block_indices=True,
)
target_token_ids = torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32, device=self.device)
target_positions = torch.tensor([0, 1, 2, 0, 1], dtype=torch.int64, device=self.device)
target_hidden_states = torch.randn(5, proposer.hidden_size, dtype=proposer.dtype, device=self.device)
next_token_ids = torch.tensor([100, 200], dtype=torch.int32, device=self.device)
num_rejected_tokens_gpu = torch.tensor([1, 0], dtype=torch.int32, device=self.device)
out_num_tokens, out_token_indices, out_cad, long_seq_args = proposer.set_inputs_first_pass(
target_token_ids=target_token_ids,
next_token_ids=next_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
token_indices_to_sample=None,
cad=common_attn_metadata,
num_rejected_tokens_gpu=num_rejected_tokens_gpu,
)
# assert function computed outputs
assert out_num_tokens == 7
expected_out_token_indices = torch.tensor([2, 6], dtype=torch.int32, device=self.device)
assert torch.equal(expected_out_token_indices, out_token_indices)
assert long_seq_args is None
# assert attrs from proposer
attrs_from_proposer: list[tuple[str, Any, Any]] = [
("input_ids", None, slice(None, out_num_tokens)),
("positions", None, slice(None, out_num_tokens)),
("is_rejected_token_mask", None, slice(None, out_num_tokens)),
("is_masked_token_mask", None, slice(None, out_num_tokens)),
]
expected_proposer = MagicMock()
expected_proposer.input_ids = torch.tensor([10, 11, 100, 0, 20, 21, 200], dtype=torch.int32, device=self.device)
expected_proposer.positions = torch.tensor([0, 1, 2, 0, 0, 1, 2], dtype=torch.int64, device=self.device)
expected_proposer.is_rejected_token_mask = torch.tensor([0, 0, 0, 1, 0, 0, 0], device=self.device, dtype=bool)
expected_proposer.is_masked_token_mask = torch.tensor([0, 0, 0, 0, 0, 0, 0], device=self.device, dtype=bool)
for attr in attrs_from_proposer:
assert_attr_equal(attr, expected_proposer, proposer)
# assert attrs from cad
attrs_from_cad: list[str | tuple[str, Any, Any]] = [
"query_start_loc_cpu",
"query_start_loc",
"seq_lens",
"num_actual_tokens",
"max_query_len",
"max_seq_len",
"slot_mapping",
"seq_lens_cpu",
"_seq_lens_cpu",
]
expected_cad = MagicMock()
expected_cad.query_start_loc_cpu = torch.tensor([0, 4, 7], dtype=torch.int32)
expected_cad.query_start_loc = expected_cad.query_start_loc_cpu.to(self.device, non_blocking=True)
expected_cad.seq_lens = torch.tensor([4, 3], device=self.device, dtype=torch.int32)
expected_cad.num_actual_tokens = 7
expected_cad.max_query_len = 4
expected_cad.max_seq_len = 4
expected_cad.slot_mapping = torch.tensor([0, 1, 2, -1, 16, 17, 18], device=self.device, dtype=torch.int64)
expected_cad.seq_lens_cpu = torch.tensor([4, 3], dtype=torch.int32)
expected_cad._seq_lens_cpu = torch.tensor([4, 3], dtype=torch.int32)
for attrition in attrs_from_cad:
assert_attr_equal(attrition, expected_cad, out_cad)
def _build_split_pcp_input_hybrid_inputs(req_scheduled_tokens: dict[str, int], hidden_size: int):
"""Build input_ids and target_hidden_states where row i is filled with i.
Using the row index as the value makes it easy to assert which original
tokens survive the per-rank slice and which slots are padding.
"""
total_tokens = sum(req_scheduled_tokens.values())
input_ids = torch.arange(total_tokens, dtype=torch.int32)
target_hidden_states = torch.arange(total_tokens, dtype=torch.float32).unsqueeze(-1).repeat(1, hidden_size)
return input_ids, target_hidden_states
# yapf: disable
@pytest.mark.parametrize(
"pcp_size, pcp_rank, req_scheduled_tokens, hidden_size,"
" expected_num_tokens, expected_input_ids, expected_hidden_first_col,"
" expected_seq_lens, expected_cu_num_tokens, expected_max_query_len",
[
# Case 1: single req, perfectly aligned to 2*pcp_size.
# ori=8, pcp=2 -> padded=8, pcp_tokens=4
# rank 0: [0,1,2,3], rank 1: [4,5,6,7]
(
2, 0, {"0": 8}, 4,
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
[4], [0, 4], 4,
),
(
2, 1, {"0": 8}, 4,
4, [4, 5, 6, 7], [4.0, 5.0, 6.0, 7.0],
[4], [0, 4], 4,
),
# Case 2: single req, needs padding.
# ori=7, pcp=2 -> padded=8, pcp_tokens=4, num_pads=1
# rank 0: [0,1,2,3] (all valid)
# rank 1: [4,5,6,PAD] -> input_ids=[4,5,6,0], hidden_first_col=[4,5,6,0]
(
2, 0, {"0": 7}, 4,
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
[4], [0, 4], 4,
),
(
2, 1, {"0": 7}, 4,
4, [4, 5, 6, 0], [4.0, 5.0, 6.0, 0.0],
[4], [0, 4], 4,
),
# Case 3: multiple reqs with different padding needs.
# req 0: ori=4 -> padded=4, pcp_tokens=2
# req 1: ori=6 -> padded=8, pcp_tokens=4, num_pads=2
# rank 0: [0,1] + [4,5,6,7]
# rank 1: [2,3] + [8,9,PAD,PAD]
(
2, 0, {"0": 4, "1": 6}, 2,
6, [0, 1, 4, 5, 6, 7], [0.0, 1.0, 4.0, 5.0, 6.0, 7.0],
[2, 4], [0, 2, 6], 4,
),
(
2, 1, {"0": 4, "1": 6}, 2,
6, [2, 3, 8, 9, 0, 0], [2.0, 3.0, 8.0, 9.0, 0.0, 0.0],
[2, 4], [0, 2, 6], 4,
),
# Case 4: pcp_size=4
# ori=9, pcp=4 -> padded=16, pcp_tokens=4, num_pads=7
# rank 0: tokens [0,1,2,3]
# rank 1: tokens [4,5,6,7]
# rank 2: tokens [8,PAD,PAD,PAD]
# rank 3: tokens [PAD,PAD,PAD,PAD]
(
4, 0, {"0": 9}, 2,
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
[4], [0, 4], 4,
),
(
4, 1, {"0": 9}, 2,
4, [4, 5, 6, 7], [4.0, 5.0, 6.0, 7.0],
[4], [0, 4], 4,
),
(
4, 2, {"0": 9}, 2,
4, [8, 0, 0, 0], [8.0, 0.0, 0.0, 0.0],
[4], [0, 4], 4,
),
(
4, 3, {"0": 9}, 2,
4, [0, 0, 0, 0], [0.0, 0.0, 0.0, 0.0],
[4], [0, 4], 4,
),
# Case 5: minimal request - single token.
# ori=1, pcp=2 -> padded=4, pcp_tokens=2, num_pads=3
# rank 0: [0, PAD]
# rank 1: [PAD, PAD]
(
2, 0, {"0": 1}, 2,
2, [0, 0], [0.0, 0.0],
[2], [0, 2], 2,
),
(
2, 1, {"0": 1}, 2,
2, [0, 0], [0.0, 0.0],
[2], [0, 2], 2,
),
],
)
# yapf: enable
def test_split_spec_decode_pcp_prefill_input_hybrid(
pcp_size,
pcp_rank,
req_scheduled_tokens,
hidden_size,
expected_num_tokens,
expected_input_ids,
expected_hidden_first_col,
expected_seq_lens,
expected_cu_num_tokens,
expected_max_query_len,
):
input_ids, target_hidden_states = _build_split_pcp_input_hybrid_inputs(
req_scheduled_tokens, hidden_size
)
# _split_spec_decode_pcp_prefill_input_hybrid only reads PCP rank/size,
# so a MagicMock with just those attributes is enough to drive the
# unbound manager method without instantiating the full manager.
mock_self = MagicMock()
mock_self.pcp_world_size = pcp_size
mock_self.pcp_world_rank = pcp_rank
(
num_tokens,
out_input_ids,
out_hidden_states,
max_query_len,
seq_lens,
cu_num_tokens,
) = PCPManager._split_spec_decode_pcp_prefill_input_hybrid(
mock_self, req_scheduled_tokens, input_ids, target_hidden_states
)
assert num_tokens == expected_num_tokens
assert max_query_len == expected_max_query_len
assert torch.equal(
out_input_ids, torch.tensor(expected_input_ids, dtype=torch.int32)
)
assert torch.equal(seq_lens, torch.tensor(expected_seq_lens, dtype=torch.int32))
assert torch.equal(
cu_num_tokens, torch.tensor(expected_cu_num_tokens, dtype=torch.int64)
)
# hidden_states shape: [num_tokens, hidden_size]
assert out_hidden_states.shape == (expected_num_tokens, hidden_size)
assert torch.equal(
out_hidden_states[:, 0],
torch.tensor(expected_hidden_first_col, dtype=torch.float32),
)
def test_split_spec_decode_pcp_prefill_input_hybrid_preserves_hidden_size():
"""Hidden states must come back with the same hidden dim as the input."""
hidden_size = 7
req_scheduled_tokens = {"0": 6}
input_ids, target_hidden_states = _build_split_pcp_input_hybrid_inputs(
req_scheduled_tokens, hidden_size
)
mock_self = MagicMock()
mock_self.pcp_world_size = 2
mock_self.pcp_world_rank = 0
_, _, out_hidden_states, _, _, _ = (
PCPManager._split_spec_decode_pcp_prefill_input_hybrid(
mock_self, req_scheduled_tokens, input_ids, target_hidden_states
)
)
assert out_hidden_states.shape[1] == hidden_size
class TestDeepSeekMTPIndicesSharing(unittest.TestCase):
"""
Unit tests for DeepSeek Sparse MLA MTP Layer's Top-K Index reuse feature (PR #10510)
"""
def setUp(self):
# Prepare the base Config Mock
self.vllm_config = MagicMock(spec=VllmConfig)
self.vllm_config.speculative_config = MagicMock()
self.vllm_config.speculative_config.draft_model_config = MagicMock()
self.vllm_config.speculative_config.draft_model_config.hf_config = MagicMock()
self.device = torch.device("cpu")
self.runner = MagicMock()
self.runner.pcp_manager = None
def test_init_mtp_indices_flag(self):
"""Test whether index_share_for_mtp_iteration is correctly read in __init__."""
# Scenario 1: Set to True in config
self.vllm_config.speculative_config.draft_model_config.hf_config.index_share_for_mtp_iteration = True
with patch.object(AscendEagleProposer, "__init__", lambda self, vllm_config, device, runner: None):
proposer = AscendEagleProposer(self.vllm_config, self.device, self.runner)
# Manually trigger the newly added initialization logic for validation
proposer.vllm_config = self.vllm_config
proposer._share_mtp_indices = getattr(
self.vllm_config.speculative_config.draft_model_config.hf_config, "index_share_for_mtp_iteration", False
)
self.assertTrue(proposer._share_mtp_indices, "MTP share flag should be True when configured.")
# Scenario 2: Not set in config (defaults to False)
del self.vllm_config.speculative_config.draft_model_config.hf_config.index_share_for_mtp_iteration
with patch.object(AscendEagleProposer, "__init__", lambda self, vllm_config, device, runner: None):
proposer2 = AscendEagleProposer(self.vllm_config, self.device, self.runner)
proposer2.vllm_config = self.vllm_config
proposer2._share_mtp_indices = getattr(
self.vllm_config.speculative_config.draft_model_config.hf_config, "index_share_for_mtp_iteration", False
)
self.assertFalse(proposer2._share_mtp_indices, "MTP share flag should default to False.")
def test_maybe_share_topk_indices_submodules(self):
"""Test if _maybe_share_topk_indices correctly updates topk_indices_buffer for all submodules."""
# Use __new__ to bypass the complex __init__ process
proposer = AscendEagleProposer.__new__(AscendEagleProposer)
# 1. Mock Target Model
target_model = MagicMock()
target_buffer_mock = MagicMock()
target_model.model.topk_indices_buffer = target_buffer_mock
# 2. Mock Draft Model (including submodules)
draft_model_mock = MagicMock()
draft_model_mock.model.topk_indices_buffer = MagicMock() # Old buffer
# Construct several submodules: some have topk_indices_buffer, some don't
mod1 = MagicMock()
mod1.topk_indices_buffer = MagicMock() # This should be replaced
mod2 = MagicMock()
del mod2.topk_indices_buffer # This doesn't have the attribute, shouldn't throw an error
mod3 = MagicMock()
mod3.topk_indices_buffer = MagicMock() # This should also be replaced
# Mock the return value of named_modules
draft_model_mock.model.named_modules.return_value = [("layer.0", mod1), ("layer.1", mod2), ("layer.2", mod3)]
proposer.model = draft_model_mock
# Execute the target method
proposer._maybe_share_topk_indices(target_model)
# Assertion: The outermost draft model buffer is updated
self.assertEqual(proposer.model.model.topk_indices_buffer, target_buffer_mock)
# Assertion: Submodules are correctly traversed and updated
self.assertEqual(mod1.topk_indices_buffer, target_buffer_mock, "Module 1 buffer should be updated.")
self.assertEqual(mod3.topk_indices_buffer, target_buffer_mock, "Module 3 buffer should be updated.")
self.assertFalse(hasattr(mod2, "topk_indices_buffer"), "Module 2 should not have a buffer added.")
def test_run_merge_draft_mtp_skip_topk(self):
"""Test the set_skip_topk calling logic in step 0 and step 1 of run_merge_draft."""
proposer = AscendEagleProposer.__new__(AscendEagleProposer)
proposer._share_mtp_indices = True
# Mock Draft Model and its set_skip_topk method
draft_model_mock = MagicMock()
proposer.model = MagicMock()
proposer.model.model = draft_model_mock
# Mock model inference return
proposer.model_returns_tuple = MagicMock(return_value=False)
proposer.model.return_value = MagicMock()
# Mock the run_merge_draft logic from your PR
# (Since this is a class method, we use an inner function to simulate and verify the core logic)
def mock_run_merge_draft(**model_kwargs):
# Step 0
draft_model = getattr(proposer.model, "model", None)
if proposer._share_mtp_indices and draft_model is not None and hasattr(draft_model, "set_skip_topk"):
draft_model.set_skip_topk(False)
# (Model inference...)
proposer.model(**model_kwargs)
# Step 1
if proposer._share_mtp_indices and draft_model is not None and hasattr(draft_model, "set_skip_topk"):
draft_model.set_skip_topk(True)
# Run the test
mock_run_merge_draft(input_ids=torch.tensor([1, 2, 3]))
# Assert calling conditions
self.assertTrue(draft_model_mock.set_skip_topk.called, "set_skip_topk should be called.")
# Verify calling order: False first, then True
calls = draft_model_mock.set_skip_topk.call_args_list
self.assertEqual(len(calls), 2, "set_skip_topk should be called exactly twice.")
self.assertEqual(calls[0][0][0], False, "Step 0 should call set_skip_topk(False).")
self.assertEqual(calls[1][0][0], True, "Step 1 should call set_skip_topk(True).")