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