# 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).")