0
tests/ut/spec_decode/__init__.py
Normal file
0
tests/ut/spec_decode/__init__.py
Normal file
0
tests/ut/spec_decode/a2/__init__.py
Normal file
0
tests/ut/spec_decode/a2/__init__.py
Normal file
4558
tests/ut/spec_decode/a2/test_eagle_proposer.py
Normal file
4558
tests/ut/spec_decode/a2/test_eagle_proposer.py
Normal file
File diff suppressed because it is too large
Load Diff
458
tests/ut/spec_decode/test_extract_hidden_states_proposer.py
Normal file
458
tests/ut/spec_decode/test_extract_hidden_states_proposer.py
Normal file
@@ -0,0 +1,458 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Unit tests for AscendExtractHiddenStatesProposer.
|
||||
|
||||
This test file follows the pattern from vllm's test_extract_hidden_states.py,
|
||||
with Ascend-specific additions for ACL graph differences.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.config import CacheConfig, CUDAGraphMode, VllmConfig, set_current_vllm_config
|
||||
|
||||
from vllm_ascend.ascend_config import init_ascend_config
|
||||
from vllm_ascend.spec_decode.extract_hidden_states_proposer import (
|
||||
AscendExtractHiddenStatesProposer,
|
||||
)
|
||||
from vllm_ascend.utils import vllm_version_is
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_pin_memory():
|
||||
if vllm_version_is("0.23.0"):
|
||||
with patch(
|
||||
"vllm.v1.spec_decode.extract_hidden_states.is_pin_memory_available",
|
||||
return_value=False,
|
||||
):
|
||||
yield
|
||||
else:
|
||||
with patch(
|
||||
"vllm.v1.spec_decode.extract_hidden_states.PIN_MEMORY",
|
||||
False,
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
class MockCachedRequestState:
|
||||
"""Mock CachedRequestState for testing (same pattern as vllm)."""
|
||||
|
||||
def __init__(self, req_id: str, token_ids: list[int]):
|
||||
self.req_id = req_id
|
||||
self.token_ids = token_ids
|
||||
|
||||
def get_token_id(self, position: int) -> int:
|
||||
if 0 <= position < len(self.token_ids):
|
||||
return self.token_ids[position]
|
||||
return 0
|
||||
|
||||
|
||||
class MockInputBatch:
|
||||
"""Mock InputBatch for testing (same pattern as vllm)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_reqs: int,
|
||||
req_ids: list[str],
|
||||
vocab_size: int,
|
||||
num_tokens_no_spec: list[int] | None = None,
|
||||
):
|
||||
self.num_reqs = num_reqs
|
||||
self.req_ids = req_ids
|
||||
self.vocab_size = vocab_size
|
||||
if num_tokens_no_spec is None:
|
||||
self.num_tokens_no_spec = np.array([5] * num_reqs, dtype=np.int64)
|
||||
else:
|
||||
self.num_tokens_no_spec = np.array(num_tokens_no_spec, dtype=np.int64)
|
||||
|
||||
|
||||
def _create_vllm_config(num_speculative_tokens: int = 1, layer_ids: list[int] | None = None):
|
||||
"""Create a VllmConfig for testing (simplified version of vllm's pattern)."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
if layer_ids is None:
|
||||
layer_ids = [1, 2, 3, 4]
|
||||
|
||||
vllm_config = MagicMock(spec=VllmConfig)
|
||||
vllm_config.speculative_config = MagicMock()
|
||||
vllm_config.speculative_config.num_speculative_tokens = num_speculative_tokens
|
||||
vllm_config.speculative_config.draft_tensor_parallel_size = 1
|
||||
vllm_config.speculative_config.draft_model_config = MagicMock()
|
||||
vllm_config.speculative_config.draft_model_config.uses_xdrope_dim = 0
|
||||
vllm_config.speculative_config.draft_model_config.uses_mrope = False
|
||||
vllm_config.speculative_config.disable_padded_drafter_batch = False
|
||||
|
||||
vllm_config.cache_config = MagicMock(spec=CacheConfig)
|
||||
vllm_config.cache_config.block_size = 16
|
||||
|
||||
vllm_config.scheduler_config = MagicMock()
|
||||
vllm_config.scheduler_config.max_num_batched_tokens = 1024
|
||||
vllm_config.scheduler_config.max_num_seqs = 32
|
||||
|
||||
vllm_config.model_config = MagicMock()
|
||||
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.model_config.hf_text_config = MagicMock(spec=[])
|
||||
vllm_config.model_config.hf_text_config.to_dict = MagicMock(return_value={})
|
||||
vllm_config.model_config.get_hidden_size = MagicMock(return_value=4096)
|
||||
|
||||
vllm_config.compilation_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 = None
|
||||
|
||||
init_ascend_config(vllm_config)
|
||||
return vllm_config
|
||||
|
||||
|
||||
def test_proposer_initialization():
|
||||
"""Test that the proposer initializes correctly (matches vllm pattern)."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
vllm_config = _create_vllm_config(num_speculative_tokens=1, layer_ids=[1, 2, 3, 4])
|
||||
device = torch.device("cpu")
|
||||
runner = MagicMock()
|
||||
runner.pin_memory = False
|
||||
runner.pcp_size = 1
|
||||
runner.dcp_size = 1
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=device, runner=runner)
|
||||
|
||||
# Verify it's an instance of ExtractHiddenStatesProposer
|
||||
from vllm.v1.spec_decode.extract_hidden_states import ExtractHiddenStatesProposer
|
||||
|
||||
assert isinstance(proposer, ExtractHiddenStatesProposer)
|
||||
assert proposer.runner == runner
|
||||
|
||||
|
||||
def test_dummy_run_basic():
|
||||
"""Test dummy_run with Ascend-specific ACL graph signature.
|
||||
|
||||
This is Ascend-specific because ACL graph capture has different parameters
|
||||
than CUDA graph capture.
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
vllm_config = _create_vllm_config()
|
||||
device = torch.device("cpu")
|
||||
runner = MagicMock()
|
||||
runner.pin_memory = False
|
||||
runner.pcp_size = 1
|
||||
runner.dcp_size = 1
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=device, runner=runner)
|
||||
|
||||
proposer.model = MagicMock()
|
||||
proposer.dp_rank = 0
|
||||
proposer.hidden_states = torch.zeros(1024, 4096, dtype=torch.float16)
|
||||
runner._sync_metadata_across_dp.return_value = (16, None, CUDAGraphMode.NONE)
|
||||
|
||||
with patch("vllm_ascend.spec_decode.extract_hidden_states_proposer.set_forward_context") as mock_context:
|
||||
mock_context.return_value.__enter__ = MagicMock(return_value=None)
|
||||
mock_context.return_value.__exit__ = MagicMock(return_value=None)
|
||||
|
||||
proposer.dummy_run(num_tokens=16)
|
||||
proposer.model.assert_called_once()
|
||||
|
||||
|
||||
def test_dummy_run_syncs_metadata_across_dp_as_draft_model():
|
||||
"""dummy_run must issue the same drafter DP sync as propose() does on
|
||||
busy ranks (via _determine_batch_execution_and_padding), mirroring
|
||||
llm_base_proposer.dummy_run.
|
||||
|
||||
Regression guard for the multi-DP deadlock: if idle DP ranks running the
|
||||
dummy path skip the drafter sync while busy ranks perform it, the DP
|
||||
cpu_group collectives desynchronize and all ranks hang.
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
vllm_config = _create_vllm_config()
|
||||
device = torch.device("cpu")
|
||||
runner = MagicMock()
|
||||
runner.pin_memory = False
|
||||
runner.pcp_size = 1
|
||||
runner.dcp_size = 1
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=device, runner=runner)
|
||||
|
||||
proposer.model = MagicMock()
|
||||
proposer.dp_rank = 0
|
||||
proposer.hidden_states = torch.zeros(1024, 4096, dtype=torch.float16)
|
||||
|
||||
synced_tensor = torch.tensor([16, 16], dtype=torch.int32)
|
||||
runner._sync_metadata_across_dp.return_value = (16, synced_tensor, CUDAGraphMode.NONE)
|
||||
|
||||
with patch("vllm_ascend.spec_decode.extract_hidden_states_proposer.set_forward_context") as mock_context:
|
||||
mock_context.return_value.__enter__ = MagicMock(return_value=None)
|
||||
mock_context.return_value.__exit__ = MagicMock(return_value=None)
|
||||
|
||||
proposer.dummy_run(num_tokens=16)
|
||||
|
||||
runner._sync_metadata_across_dp.assert_called_once()
|
||||
args, kwargs = runner._sync_metadata_across_dp.call_args
|
||||
assert (args and args[0] == 16) or kwargs.get("num_tokens") == 16
|
||||
assert kwargs.get("is_draft_model") is True
|
||||
# The synced tensor must be the one forwarded to set_forward_context.
|
||||
_, ctx_kwargs = mock_context.call_args
|
||||
assert ctx_kwargs["num_tokens_across_dp"] is synced_tensor
|
||||
|
||||
|
||||
def test_prepare_next_token_ids_padded():
|
||||
"""Test prepare_next_token_ids_padded (matches vllm's test pattern).
|
||||
|
||||
Since num_speculative_tokens == 1, sampled_token_ids has shape (batch_size, 1).
|
||||
For each request we either use the sampled token (if valid and not discarded)
|
||||
or a backup token from the request state.
|
||||
|
||||
Note: Ascend uses indices/count pattern instead of GPU's boolean mask.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
device = torch.device("cpu")
|
||||
vllm_config = _create_vllm_config()
|
||||
|
||||
runner = MagicMock()
|
||||
runner.pin_memory = False
|
||||
runner.pcp_size = 1
|
||||
runner.dcp_size = 1
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=device, runner=runner)
|
||||
|
||||
# Setup test data (same pattern as vllm)
|
||||
num_requests = 4
|
||||
req_ids = [f"req_{i + 1}" for i in range(num_requests)]
|
||||
|
||||
gpu_input_batch = MockInputBatch(
|
||||
num_reqs=num_requests,
|
||||
req_ids=req_ids,
|
||||
vocab_size=100,
|
||||
num_tokens_no_spec=[11, 16, 21, 26], # Different seq_lens: [10, 15, 20, 25]
|
||||
)
|
||||
|
||||
requests = {}
|
||||
for req_id in req_ids:
|
||||
idx = int(req_id.split("_")[1])
|
||||
# Different token sequences for each request
|
||||
mock_request = MockCachedRequestState(req_id, list(range(15 + idx * 5)))
|
||||
requests[req_id] = mock_request
|
||||
|
||||
# sampled_token_ids shape: [batch_size, 1]
|
||||
sampled_token_ids = torch.tensor(
|
||||
[
|
||||
[1], # valid, use 1
|
||||
[4], # valid, use 4
|
||||
[-1], # invalid, use backup
|
||||
[2], # discarded, use backup
|
||||
],
|
||||
dtype=torch.int64,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# Ascend uses indices/count pattern (different from GPU's boolean mask)
|
||||
discard_request_indices = torch.tensor([3], dtype=torch.int64)
|
||||
num_discarded_requests = 1
|
||||
|
||||
next_token_ids, valid_sampled_tokens_count = 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 results
|
||||
# valid_sampled_tokens_count tracks token validity (not discard status)
|
||||
expected_valid_counts = torch.tensor([1, 1, 0, 1], dtype=torch.int32)
|
||||
assert torch.equal(valid_sampled_tokens_count, expected_valid_counts)
|
||||
|
||||
# next_token_ids: use sampled if valid and not discarded, else backup
|
||||
# Request 1: valid (1), use 1
|
||||
# Request 2: valid (4), use 4
|
||||
# Request 3: invalid (-1), use backup (seq_len=10, token at pos 10 = 10)
|
||||
# Request 4: discarded, use backup (seq_len=15, token at pos 15 = 15)
|
||||
assert next_token_ids[0].item() == 1
|
||||
assert next_token_ids[1].item() == 4
|
||||
assert next_token_ids[2].item() == 20 # backup from request state
|
||||
assert next_token_ids[3].item() == 25 # backup from request state
|
||||
|
||||
# Verify return dtypes
|
||||
assert next_token_ids.dtype == torch.int32
|
||||
assert valid_sampled_tokens_count.dtype == torch.int32
|
||||
|
||||
|
||||
def _make_batch_desc(num_tokens: int):
|
||||
"""Build a minimal mock for ``CudagraphDispatcher.dispatch``'s batch_desc."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
batch_desc = MagicMock()
|
||||
batch_desc.num_tokens = num_tokens
|
||||
return batch_desc
|
||||
|
||||
|
||||
def _build_proposer_for_padding_test(data_parallel_size: int = 1):
|
||||
"""Shared helper for _determine_batch_execution_and_padding tests.
|
||||
|
||||
Returns a proposer whose ``runner``, ``cudagraph_dispatcher``, and
|
||||
DP-related attributes are mocked so we can drive
|
||||
``_determine_batch_execution_and_padding`` without an NPU.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
vllm_config = _create_vllm_config()
|
||||
vllm_config.parallel_config.data_parallel_size = data_parallel_size
|
||||
|
||||
runner = MagicMock()
|
||||
runner.pin_memory = False
|
||||
runner.pcp_size = 1
|
||||
runner.dcp_size = 1
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=torch.device("cpu"), runner=runner)
|
||||
|
||||
proposer.dp_rank = 0
|
||||
proposer.cudagraph_dispatcher = MagicMock()
|
||||
return proposer, runner
|
||||
|
||||
|
||||
def test_determine_batch_execution_and_padding_asserts_when_runner_is_none():
|
||||
"""Constructing without a runner must fail fast with a clear message.
|
||||
|
||||
Regression guard for the AttributeError that would otherwise be raised
|
||||
on ``self.runner._pad_for_sequence_parallelism(...)`` at the entry of
|
||||
the override.
|
||||
"""
|
||||
vllm_config = _create_vllm_config()
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
proposer = AscendExtractHiddenStatesProposer(vllm_config=vllm_config, device=torch.device("cpu"), runner=None)
|
||||
|
||||
proposer.cudagraph_dispatcher = type("D", (), {"dispatch": staticmethod(lambda *a, **kw: (None, None))})()
|
||||
|
||||
with pytest.raises(AssertionError, match="requires a runner reference"):
|
||||
proposer._determine_batch_execution_and_padding(num_tokens=4)
|
||||
|
||||
|
||||
def test_determine_batch_execution_and_padding_dp1_sp_pads_and_skips_sync():
|
||||
"""With DP=1, SP-pads ``num_tokens`` but never calls DP sync.
|
||||
|
||||
Verifies the ``data_parallel_size == 1`` early-out and that the
|
||||
runner's ``_pad_for_sequence_parallelism`` is still consulted so the
|
||||
cache_only forward gets an SP-aligned input.
|
||||
"""
|
||||
proposer, runner = _build_proposer_for_padding_test(data_parallel_size=1)
|
||||
|
||||
# Simulate TP=4 SP padding: round 6 up to 8
|
||||
runner._pad_for_sequence_parallelism = lambda n: ((n + 3) // 4) * 4
|
||||
proposer.cudagraph_dispatcher.dispatch.return_value = (
|
||||
CUDAGraphMode.NONE,
|
||||
_make_batch_desc(num_tokens=8),
|
||||
)
|
||||
|
||||
cudagraph_mode, num_tokens_padded, num_tokens_across_dp = proposer._determine_batch_execution_and_padding(
|
||||
num_tokens=6
|
||||
)
|
||||
|
||||
assert num_tokens_padded == 8
|
||||
assert num_tokens_across_dp is None # no DP sync when dp_size == 1
|
||||
assert cudagraph_mode == CUDAGraphMode.NONE
|
||||
# Dispatcher saw the SP-padded value, not the raw 6.
|
||||
args, kwargs = proposer.cudagraph_dispatcher.dispatch.call_args
|
||||
assert args[0] == 8 or kwargs.get("num_tokens") == 8 or args == (8,)
|
||||
runner._sync_metadata_across_dp.assert_not_called()
|
||||
|
||||
|
||||
def test_determine_batch_execution_and_padding_dp2_uses_runner_sync():
|
||||
"""With DP>1, must call ``runner._sync_metadata_across_dp`` (Ascend's
|
||||
shape ``[2, dp_size]`` path) and must NOT call upstream
|
||||
``coordinate_batch_across_dp`` (shape ``[4, dp_size]``).
|
||||
|
||||
This is the core regression test for the gloo
|
||||
``op.preamble.length 8 vs 4`` shape mismatch on the DP cpu_group.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
proposer, runner = _build_proposer_for_padding_test(data_parallel_size=2)
|
||||
|
||||
runner._pad_for_sequence_parallelism = lambda n: ((n + 3) // 4) * 4
|
||||
proposer.cudagraph_dispatcher.dispatch.side_effect = [
|
||||
# First dispatch (pre-sync) with SP-padded num_tokens=8
|
||||
(CUDAGraphMode.NONE, _make_batch_desc(num_tokens=8)),
|
||||
# Re-dispatch after sync; the agreed value happens to also be 8
|
||||
(CUDAGraphMode.NONE, _make_batch_desc(num_tokens=8)),
|
||||
]
|
||||
# Pretend both DP ranks agreed on 8 tokens
|
||||
sync_tensor = torch.tensor([8, 8], dtype=torch.int32)
|
||||
runner._sync_metadata_across_dp.return_value = (8, sync_tensor, CUDAGraphMode.NONE)
|
||||
|
||||
with patch("vllm.v1.spec_decode.extract_hidden_states.coordinate_batch_across_dp") as mock_upstream_coord:
|
||||
cudagraph_mode, num_tokens_padded, num_tokens_across_dp = proposer._determine_batch_execution_and_padding(
|
||||
num_tokens=6
|
||||
)
|
||||
|
||||
# Upstream DP sync must NOT be used (it would post a [4, dp_size]
|
||||
# tensor and break gloo on the shared cpu_group).
|
||||
mock_upstream_coord.assert_not_called()
|
||||
# Runner sync called once, with the SP-padded value and is_draft_model=True.
|
||||
runner._sync_metadata_across_dp.assert_called_once()
|
||||
call_kwargs = runner._sync_metadata_across_dp.call_args.kwargs
|
||||
assert call_kwargs["num_tokens"] == 8 # SP-padded 6 -> 8
|
||||
assert call_kwargs["is_draft_model"] is True
|
||||
|
||||
assert num_tokens_padded == 8
|
||||
assert num_tokens_across_dp is not None
|
||||
assert num_tokens_across_dp[proposer.dp_rank].item() == 8
|
||||
|
||||
|
||||
def test_determine_batch_execution_and_padding_dp2_keeps_tp_aligned_for_main_forward():
|
||||
"""If the runner's SP padding produces a TP-aligned value, the final
|
||||
``num_tokens_padded`` returned to the proposer (and downstream main
|
||||
forward) is guaranteed to be TP-aligned too. Regression guard for
|
||||
the ``reduce_scatter`` assertion ``input.shape[0] % world_size == 0``.
|
||||
"""
|
||||
proposer, runner = _build_proposer_for_padding_test(data_parallel_size=2)
|
||||
|
||||
tp = 4
|
||||
runner._pad_for_sequence_parallelism = lambda n: ((n + tp - 1) // tp) * tp
|
||||
proposer.cudagraph_dispatcher.dispatch.side_effect = [
|
||||
(CUDAGraphMode.NONE, _make_batch_desc(num_tokens=8)),
|
||||
(CUDAGraphMode.NONE, _make_batch_desc(num_tokens=8)),
|
||||
]
|
||||
runner._sync_metadata_across_dp.return_value = (
|
||||
8,
|
||||
torch.tensor([8, 8], dtype=torch.int32),
|
||||
CUDAGraphMode.NONE,
|
||||
)
|
||||
|
||||
_mode, num_tokens_padded, _across = proposer._determine_batch_execution_and_padding(num_tokens=6)
|
||||
|
||||
# The whole point of the fix: never returns 6 (which would crash
|
||||
# SP reduce_scatter as 6 % 4 != 0).
|
||||
assert num_tokens_padded % tp == 0
|
||||
112
tests/ut/spec_decode/test_llm_base_proposer.py
Normal file
112
tests/ut/spec_decode/test_llm_base_proposer.py
Normal file
@@ -0,0 +1,112 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from vllm.config import CUDAGraphMode
|
||||
|
||||
from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer
|
||||
|
||||
# CUDAGraphMode values whose ``has_full_cudagraphs()`` is True: FULL plus the
|
||||
# two composite modes that mix FULL with NONE / PIECEWISE.
|
||||
FULL_CUDAGRAPH_MODES = [
|
||||
CUDAGraphMode.FULL,
|
||||
CUDAGraphMode.FULL_DECODE_ONLY,
|
||||
CUDAGraphMode.FULL_AND_PIECEWISE,
|
||||
]
|
||||
|
||||
# Modes without a full cudagraph.
|
||||
NON_FULL_CUDAGRAPH_MODES = [
|
||||
CUDAGraphMode.NONE,
|
||||
CUDAGraphMode.PIECEWISE,
|
||||
]
|
||||
|
||||
|
||||
class TestDisablePaddedDrafterBatchWithFullGraph:
|
||||
"""Guard: ``disable_padded_drafter_batch=True`` + cuda graph + any full
|
||||
cudagraph mode must raise ``NotImplementedError``.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _make_proposer(
|
||||
*,
|
||||
disable_padded_drafter_batch: bool,
|
||||
use_cuda_graph: bool,
|
||||
cudagraph_mode: CUDAGraphMode,
|
||||
) -> AscendSpecDecodeBaseProposer:
|
||||
"""Bypass ``__init__`` and set only the three attrs the guard reads.
|
||||
|
||||
``cudagraph_mode`` is a real enum value so ``has_full_cudagraphs()`` is
|
||||
exercised, not stubbed.
|
||||
"""
|
||||
proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer)
|
||||
proposer.speculative_config = SimpleNamespace(
|
||||
disable_padded_drafter_batch=disable_padded_drafter_batch,
|
||||
)
|
||||
proposer.use_cuda_graph = use_cuda_graph
|
||||
proposer.compilation_config = SimpleNamespace(cudagraph_mode=cudagraph_mode)
|
||||
return proposer
|
||||
|
||||
@pytest.mark.parametrize("cudagraph_mode", FULL_CUDAGRAPH_MODES)
|
||||
def test_guard_raises_when_padded_drafter_batch_disabled_with_full_cudagraph(self, cudagraph_mode: CUDAGraphMode):
|
||||
"""The bad combo: disable_padded + cuda graph + any full-cudagraph mode
|
||||
is intercepted with ``NotImplementedError``."""
|
||||
proposer = self._make_proposer(
|
||||
disable_padded_drafter_batch=True,
|
||||
use_cuda_graph=True,
|
||||
cudagraph_mode=cudagraph_mode,
|
||||
)
|
||||
|
||||
with pytest.raises(NotImplementedError, match="disable_padded_drafter_batch"):
|
||||
proposer._raise_if_padded_drafter_batch_disabled_and_full_graph_enabled()
|
||||
|
||||
@pytest.mark.parametrize("cudagraph_mode", NON_FULL_CUDAGRAPH_MODES)
|
||||
def test_guard_does_not_raise_without_full_cudagraph(self, cudagraph_mode: CUDAGraphMode):
|
||||
"""NONE / PIECEWISE never trip the guard, even with disable_padded + cuda graph."""
|
||||
proposer = self._make_proposer(
|
||||
disable_padded_drafter_batch=True,
|
||||
use_cuda_graph=True,
|
||||
cudagraph_mode=cudagraph_mode,
|
||||
)
|
||||
|
||||
# Must not raise.
|
||||
proposer._raise_if_padded_drafter_batch_disabled_and_full_graph_enabled()
|
||||
|
||||
@pytest.mark.parametrize("cudagraph_mode", FULL_CUDAGRAPH_MODES)
|
||||
def test_guard_does_not_raise_when_padded_drafter_batch_enabled(self, cudagraph_mode: CUDAGraphMode):
|
||||
"""Padded drafter batch on (the default) is fine with any full cudagraph."""
|
||||
proposer = self._make_proposer(
|
||||
disable_padded_drafter_batch=False,
|
||||
use_cuda_graph=True,
|
||||
cudagraph_mode=cudagraph_mode,
|
||||
)
|
||||
|
||||
proposer._raise_if_padded_drafter_batch_disabled_and_full_graph_enabled()
|
||||
|
||||
def test_guard_does_not_raise_when_eager(self):
|
||||
"""``enforce_eager`` -> ``use_cuda_graph=False`` short-circuits the guard."""
|
||||
proposer = self._make_proposer(
|
||||
disable_padded_drafter_batch=True,
|
||||
use_cuda_graph=False,
|
||||
cudagraph_mode=CUDAGraphMode.FULL,
|
||||
)
|
||||
|
||||
proposer._raise_if_padded_drafter_batch_disabled_and_full_graph_enabled()
|
||||
424
tests/ut/spec_decode/test_speculators_vwn_eagle3.py
Normal file
424
tests/ut/spec_decode/test_speculators_vwn_eagle3.py
Normal file
@@ -0,0 +1,424 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Unit tests for VWN-Eagle3 model components.
|
||||
|
||||
Tests cover PreVwnLayerV1, VwnLlamaDecoderLayer, VwnLlamaModel, and
|
||||
Eagle3VwnLlamaForCausalLM using CPU-only execution with mocked VllmConfig.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from vllm.config import CacheConfig, CompilationMode, VllmConfig, set_current_vllm_config
|
||||
|
||||
from vllm_ascend.ascend_config import init_ascend_config
|
||||
from vllm_ascend.models.llama_eagle3_vwn import (
|
||||
Eagle3VwnLlamaForCausalLM,
|
||||
PreVwnLayerV1,
|
||||
VwnLlamaDecoderLayer,
|
||||
VwnLlamaModel,
|
||||
)
|
||||
|
||||
_HIDDEN = 2048
|
||||
_INTERMEDIATE = 6144
|
||||
_VOCAB = 151936
|
||||
_DRAFT_VOCAB = 35000
|
||||
_NUM_HEADS = 32
|
||||
_NUM_KV_HEADS = 4
|
||||
_RMS_EPS = 1e-6
|
||||
|
||||
|
||||
class _PassthroughAttn(nn.Module):
|
||||
"""Replaces self_attn for CPU tests — returns input unchanged."""
|
||||
|
||||
def forward(self, *, positions, hidden_states, **kwargs):
|
||||
return hidden_states
|
||||
|
||||
|
||||
class _PassthroughMLP(nn.Module):
|
||||
"""Replaces mlp for CPU tests — returns input unchanged."""
|
||||
|
||||
def forward(self, hidden_states):
|
||||
return hidden_states
|
||||
|
||||
|
||||
class _MockTPGroup:
|
||||
"""Minimal mock for get_tp_group() when TP=1."""
|
||||
|
||||
rank_in_group = 0
|
||||
world_size = 1
|
||||
|
||||
def all_reduce(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def all_gather(self, x, *args, **kwargs):
|
||||
return x.unsqueeze(0)
|
||||
|
||||
def reduce_scatter(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
|
||||
def _mock_npu_ops_on_layer(layer):
|
||||
"""Replace self_attn and mlp with passthrough modules for CPU testing."""
|
||||
layer.self_attn = _PassthroughAttn()
|
||||
layer.mlp = _PassthroughMLP()
|
||||
|
||||
|
||||
def _cpu_rms_norm(x, weight, eps):
|
||||
"""CPU fallback for torch_npu.npu_rms_norm.
|
||||
|
||||
Returns the normalized tensor and a placeholder rstd (None), matching the
|
||||
2-tuple shape the production op yields so callers can unpack it.
|
||||
"""
|
||||
orig_dtype = x.dtype
|
||||
x32 = x.float()
|
||||
var = x32.pow(2).mean(-1, keepdim=True)
|
||||
x32 = x32 * torch.rsqrt(var + eps)
|
||||
out = (x32 * weight.float()).to(orig_dtype)
|
||||
return out, None
|
||||
|
||||
|
||||
def _cpu_add_rms_norm(x, residual, weight, eps):
|
||||
"""CPU fallback for torch_npu.npu_add_rms_norm (returns 3-tuple)."""
|
||||
x_plus_res = x + residual
|
||||
out, _ = _cpu_rms_norm(x_plus_res, weight, eps)
|
||||
return out, None, x_plus_res
|
||||
|
||||
|
||||
def _cpu_add_rms_norm_bias(x, residual, weight, bias, eps):
|
||||
"""CPU fallback for torch.ops._C_ascend.npu_add_rms_norm_bias."""
|
||||
out, _, new_residual = _cpu_add_rms_norm(x, residual, weight, eps)
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
return out, _, new_residual
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _mock_npu_env():
|
||||
"""Patch TP group, Ascend config, and NPU ops so all tests run on CPU.
|
||||
|
||||
conftest.py stubs ``torch_npu.npu_rms_norm`` with a bare ``MagicMock()``,
|
||||
which yields an empty iterator and breaks tuple unpacking. We override it
|
||||
here (plus the related add/bias variants and the weight-prefetch hook) with
|
||||
pure-torch CPU implementations so AscendRMSNorm can run on CPU runners.
|
||||
"""
|
||||
import torch_npu
|
||||
|
||||
_mock = _MockTPGroup()
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.enable_flashcomm2_parallel_size = 0
|
||||
mock_cfg.enable_context_parallel = False
|
||||
mock_cfg.enable_flashcomm1 = False
|
||||
mock_cfg.enable_matmul_allreduce = False
|
||||
mock_cfg.weight_nz_mode = 1
|
||||
mock_cfg.enable_mlapo = True
|
||||
mock_cfg.enable_fused_mc2 = 0
|
||||
mock_cfg.msmonitor_use_daemon = False
|
||||
mock_cfg.enable_transpose_kv_cache_by_block = True
|
||||
mock_cfg.finegrained_tp_config = MagicMock(
|
||||
lmhead_tensor_parallel_size=0,
|
||||
embedding_tensor_parallel_size=0,
|
||||
oproj_tensor_parallel_size=0,
|
||||
olora_tensor_parallel_size=0,
|
||||
mlp_tensor_parallel_size=0,
|
||||
)
|
||||
|
||||
_prefetch_mock = MagicMock()
|
||||
|
||||
with (
|
||||
patch("vllm_ascend.ops.linear_op.get_tp_group", return_value=_mock),
|
||||
patch("vllm.distributed.parallel_state.get_tp_group", return_value=_mock),
|
||||
patch("vllm_ascend.ops.vocab_parallel_embedding.get_tp_group", return_value=_mock),
|
||||
patch("vllm_ascend.utils.get_ascend_config", return_value=mock_cfg),
|
||||
patch.object(torch.ops.vllm, "unquantized_gemm", F.linear),
|
||||
patch.object(torch.ops.vllm, "maybe_calc_kv_scales", lambda *a, **kw: None),
|
||||
patch.object(torch.ops.vllm, "maybe_pad_and_reduce", lambda x, *a, **kw: x),
|
||||
patch("vllm.model_executor.layers.logits_processor.tensor_model_parallel_all_gather", lambda x, *a, **kw: x),
|
||||
patch.object(torch_npu, "npu_rms_norm", side_effect=_cpu_rms_norm, create=True),
|
||||
patch.object(torch_npu, "npu_add_rms_norm", side_effect=_cpu_add_rms_norm, create=True),
|
||||
patch.object(
|
||||
torch.ops._C_ascend,
|
||||
"npu_add_rms_norm_bias",
|
||||
side_effect=_cpu_add_rms_norm_bias,
|
||||
create=True,
|
||||
),
|
||||
patch("vllm_ascend.ops.layernorm.get_weight_prefetch_method", return_value=_prefetch_mock),
|
||||
# enable_cp() reads parallel_config.*_context_parallel_size and runs `> 1`.
|
||||
# On MagicMock these fields yield TypeError on Python 3.12, so short-circuit
|
||||
# the check everywhere it's imported.
|
||||
patch("vllm_ascend.attention.attention_v1.enable_cp", return_value=False),
|
||||
patch("vllm_ascend.attention.sfa_v1.enable_cp", return_value=False, create=True),
|
||||
patch("vllm_ascend.attention.mla_v1.enable_cp", return_value=False, create=True),
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
def _make_hf_config(
|
||||
hidden_size=_HIDDEN,
|
||||
vwn_m=4,
|
||||
vwn_r=1.5,
|
||||
num_hidden_layers=1,
|
||||
draft_vocab_size=_DRAFT_VOCAB,
|
||||
**extra,
|
||||
):
|
||||
"""Create a real LlamaConfig with VWN attributes.
|
||||
|
||||
Using a real config object avoids whack-a-mole with missing attributes
|
||||
that LlamaDecoderLayer's deep init chain expects.
|
||||
"""
|
||||
from transformers import LlamaConfig
|
||||
|
||||
cfg = LlamaConfig(
|
||||
hidden_size=hidden_size,
|
||||
intermediate_size=_INTERMEDIATE,
|
||||
num_attention_heads=_NUM_HEADS,
|
||||
num_key_value_heads=_NUM_KV_HEADS,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
vocab_size=_VOCAB,
|
||||
rms_norm_eps=_RMS_EPS,
|
||||
max_position_embeddings=40960,
|
||||
)
|
||||
cfg.vwn_m = vwn_m
|
||||
cfg.vwn_r = vwn_r
|
||||
cfg.draft_vocab_size = draft_vocab_size
|
||||
for k, v in extra.items():
|
||||
setattr(cfg, k, v)
|
||||
return cfg
|
||||
|
||||
|
||||
def _create_vllm_config_for_vwn(
|
||||
vwn_m=4,
|
||||
vwn_r=1.5,
|
||||
hidden_size=_HIDDEN,
|
||||
num_hidden_layers=1,
|
||||
num_target_layers=48,
|
||||
):
|
||||
"""Create a mocked VllmConfig for VWN model instantiation on CPU."""
|
||||
hf_config = _make_hf_config(
|
||||
hidden_size=hidden_size,
|
||||
vwn_m=vwn_m,
|
||||
vwn_r=vwn_r,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
)
|
||||
|
||||
vllm_config = MagicMock(spec=VllmConfig)
|
||||
|
||||
# speculative_config
|
||||
vllm_config.speculative_config = MagicMock()
|
||||
vllm_config.speculative_config.num_speculative_tokens = 3
|
||||
vllm_config.speculative_config.draft_tensor_parallel_size = 1
|
||||
vllm_config.speculative_config.parallel_drafting = False
|
||||
vllm_config.speculative_config.disable_padded_drafter_batch = False
|
||||
vllm_config.speculative_config.draft_model_config = MagicMock(
|
||||
hf_config=hf_config,
|
||||
uses_mrope=False,
|
||||
uses_xdrope_dim=0,
|
||||
quantization=None,
|
||||
load_config=MagicMock(),
|
||||
get_hidden_size=MagicMock(return_value=hidden_size),
|
||||
get_inputs_embeds_size=MagicMock(return_value=hidden_size),
|
||||
)
|
||||
|
||||
# cache_config
|
||||
vllm_config.cache_config = MagicMock(spec=CacheConfig)
|
||||
vllm_config.cache_config.block_size = 16
|
||||
vllm_config.cache_config.kv_cache_dtype_skip_layers = None
|
||||
vllm_config.cache_config.cache_dtype = "auto"
|
||||
|
||||
# scheduler_config
|
||||
vllm_config.scheduler_config = MagicMock()
|
||||
vllm_config.scheduler_config.max_num_batched_tokens = 1024
|
||||
vllm_config.scheduler_config.max_num_seqs = 32
|
||||
|
||||
# model_config
|
||||
vllm_config.model_config = MagicMock()
|
||||
vllm_config.model_config.dtype = torch.float32
|
||||
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.model_config.enforce_eager = True
|
||||
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.hf_config = hf_config
|
||||
vllm_config.model_config.get_num_layers = MagicMock(return_value=num_target_layers)
|
||||
|
||||
# compilation_config
|
||||
vllm_config.compilation_config = MagicMock()
|
||||
vllm_config.compilation_config.mode = CompilationMode.NONE
|
||||
vllm_config.compilation_config.pass_config = MagicMock(enable_sp=False)
|
||||
vllm_config.compilation_config.custom_ops = ["none"]
|
||||
|
||||
# parallel_config
|
||||
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.decode_context_parallel_size = 1
|
||||
vllm_config.parallel_config.enable_expert_parallel = False
|
||||
|
||||
vllm_config.additional_config = None
|
||||
|
||||
init_ascend_config(vllm_config)
|
||||
return vllm_config
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _make_model_with_mocked_ops(**kwargs):
|
||||
"""Create Eagle3VwnLlamaForCausalLM with mocked attention/MLP for CPU."""
|
||||
vllm_config = _create_vllm_config_for_vwn(**kwargs)
|
||||
hs = vllm_config.speculative_config.draft_model_config.hf_config.hidden_size
|
||||
with set_current_vllm_config(vllm_config):
|
||||
model = Eagle3VwnLlamaForCausalLM(vllm_config=vllm_config, prefix="")
|
||||
for layer in model.model.layers:
|
||||
_mock_npu_ops_on_layer(layer)
|
||||
yield model, vllm_config, hs
|
||||
|
||||
|
||||
class TestPreVwnLayerV1:
|
||||
@pytest.mark.parametrize("vwn_m,vwn_r", [(4, 1.5), (1, 1.0)])
|
||||
def test_init_and_forward(self, vwn_m, vwn_r):
|
||||
"""Verify layer init and forward output shape."""
|
||||
vllm_config = _create_vllm_config_for_vwn(vwn_m=vwn_m, vwn_r=vwn_r)
|
||||
hs, batch = _HIDDEN, 4
|
||||
wd = int(hs * vwn_r)
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
layer = PreVwnLayerV1(
|
||||
vllm_config=vllm_config,
|
||||
prefix="test_prevwn",
|
||||
config=vllm_config.speculative_config.draft_model_config.hf_config,
|
||||
)
|
||||
assert layer.wider_dim == wd
|
||||
out = layer(torch.randn(batch, hs), torch.randn(batch, hs))
|
||||
|
||||
assert out.shape == (batch, wd)
|
||||
|
||||
|
||||
class TestVwnLlamaDecoderLayer:
|
||||
@pytest.mark.parametrize("vwn_m,vwn_r", [(4, 1.5), (4, 1.0)])
|
||||
def test_forward_layer0(self, vwn_m, vwn_r):
|
||||
"""VWN forward with various m/r configs — init + shape check."""
|
||||
vllm_config = _create_vllm_config_for_vwn(vwn_m=vwn_m, vwn_r=vwn_r)
|
||||
hs, batch = _HIDDEN, 4
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
layer = VwnLlamaDecoderLayer(
|
||||
vllm_config=vllm_config,
|
||||
prefix="model.layers.48",
|
||||
config=vllm_config.speculative_config.draft_model_config.hf_config,
|
||||
layer_idx=0,
|
||||
)
|
||||
_mock_npu_ops_on_layer(layer)
|
||||
out_hidden, _ = layer(
|
||||
torch.arange(batch, dtype=torch.long),
|
||||
torch.randn(batch, hs),
|
||||
torch.randn(batch, hs),
|
||||
None,
|
||||
)
|
||||
|
||||
assert out_hidden.shape == (batch, hs)
|
||||
|
||||
def test_qkv_proj_input_size_layer0(self):
|
||||
"""VWN layer 0 qkv_proj input is hidden_size (not 2*hidden_size)."""
|
||||
vllm_config = _create_vllm_config_for_vwn()
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
layer = VwnLlamaDecoderLayer(
|
||||
vllm_config=vllm_config,
|
||||
prefix="model.layers.48",
|
||||
config=vllm_config.speculative_config.draft_model_config.hf_config,
|
||||
layer_idx=0,
|
||||
)
|
||||
|
||||
assert layer.self_attn.qkv_proj.input_size == _HIDDEN
|
||||
|
||||
|
||||
class TestVwnLlamaModel:
|
||||
@pytest.mark.parametrize(
|
||||
"num_hidden_layers,use_input_embeds",
|
||||
[
|
||||
(1, False),
|
||||
(1, True),
|
||||
],
|
||||
)
|
||||
def test_forward(self, num_hidden_layers, use_input_embeds):
|
||||
"""Verify layer count, type, and forward output shapes."""
|
||||
vllm_config = _create_vllm_config_for_vwn(num_hidden_layers=num_hidden_layers)
|
||||
hs, num_tokens = _HIDDEN, 4
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
model = VwnLlamaModel(
|
||||
vllm_config=vllm_config,
|
||||
prefix="model",
|
||||
start_layer_id=48,
|
||||
)
|
||||
|
||||
assert len(model.layers) == num_hidden_layers
|
||||
for i, layer in enumerate(model.layers):
|
||||
assert isinstance(layer, VwnLlamaDecoderLayer)
|
||||
assert layer.layer_idx == i
|
||||
_mock_npu_ops_on_layer(layer)
|
||||
|
||||
input_ids = torch.randint(0, _VOCAB, (num_tokens,))
|
||||
positions = torch.arange(num_tokens, dtype=torch.long)
|
||||
hidden_states = torch.randn(num_tokens, hs)
|
||||
input_embeds = torch.randn(num_tokens, hs) if use_input_embeds else None
|
||||
|
||||
postnorm, prenorm = model(
|
||||
input_ids,
|
||||
positions,
|
||||
hidden_states,
|
||||
input_embeds=input_embeds,
|
||||
)
|
||||
|
||||
assert postnorm.shape == (num_tokens, hs)
|
||||
assert prenorm.shape == (num_tokens, hs)
|
||||
|
||||
|
||||
class TestEagle3VwnLlamaForCausalLM:
|
||||
def test_init_and_forward(self):
|
||||
with _make_model_with_mocked_ops(vwn_m=4) as (model, _, hs):
|
||||
assert isinstance(model.model, VwnLlamaModel)
|
||||
num_tokens = 3
|
||||
|
||||
input_ids = torch.randint(0, _VOCAB, (num_tokens,))
|
||||
positions = torch.arange(num_tokens, dtype=torch.long)
|
||||
|
||||
postnorm, prenorm = model(
|
||||
input_ids,
|
||||
positions,
|
||||
torch.randn(num_tokens, hs),
|
||||
)
|
||||
|
||||
assert postnorm.shape == (num_tokens, hs)
|
||||
assert prenorm.shape == (num_tokens, hs)
|
||||
|
||||
def test_embed_input_ids(self):
|
||||
vllm_config = _create_vllm_config_for_vwn()
|
||||
num_tokens = 3
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
model = Eagle3VwnLlamaForCausalLM(vllm_config=vllm_config, prefix="")
|
||||
embeds = model.embed_input_ids(torch.randint(0, _VOCAB, (num_tokens,)))
|
||||
|
||||
assert embeds.shape == (num_tokens, _HIDDEN)
|
||||
103
tests/ut/spec_decode/test_step3p5_source_regression.py
Normal file
103
tests/ut/spec_decode/test_step3p5_source_regression.py
Normal file
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Source-level regressions for Step3.5 MTP Ascend glue.
|
||||
|
||||
Importing the Step3.5 proposer can initialize runtime/device state in this
|
||||
branch. Keep these checks focused on cross-file contracts that are hard to
|
||||
exercise in a lightweight unit test, and avoid pinning the exact implementation
|
||||
sequence inside the proposer.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[3]
|
||||
STEP3P5 = ROOT / "vllm_ascend" / "spec_decode" / "step3p5.py"
|
||||
BASE_PROPOSER = ROOT / "vllm_ascend" / "spec_decode" / "llm_base_proposer.py"
|
||||
PATCH_SPEC_CFG = ROOT / "vllm_ascend" / "patch" / "platform" / "patch_speculative_config.py"
|
||||
WORKER_PATCH_INIT = ROOT / "vllm_ascend" / "patch" / "worker" / "__init__.py"
|
||||
LEGACY_STEP3P7_PATCH = ROOT / "vllm_ascend" / "patch" / "worker" / "patch_step3p5_mtp.py"
|
||||
|
||||
|
||||
def _tree(path: Path) -> ast.Module:
|
||||
return ast.parse(path.read_text())
|
||||
|
||||
|
||||
def _class(path: Path, name: str) -> ast.ClassDef:
|
||||
for node in _tree(path).body:
|
||||
if isinstance(node, ast.ClassDef) and node.name == name:
|
||||
return node
|
||||
raise AssertionError(f"class {name} not found in {path}")
|
||||
|
||||
|
||||
def _method(path: Path, cls_name: str, method_name: str) -> ast.FunctionDef:
|
||||
cls = _class(path, cls_name)
|
||||
for node in cls.body:
|
||||
if isinstance(node, ast.FunctionDef) and node.name == method_name:
|
||||
return node
|
||||
raise AssertionError(f"method {cls_name}.{method_name} not found")
|
||||
|
||||
|
||||
def _func(path: Path, name: str) -> ast.FunctionDef:
|
||||
for node in _tree(path).body:
|
||||
if isinstance(node, ast.FunctionDef) and node.name == name:
|
||||
return node
|
||||
raise AssertionError(f"function {name} not found in {path}")
|
||||
|
||||
|
||||
def _src(node: ast.AST) -> str:
|
||||
return ast.unparse(node)
|
||||
|
||||
|
||||
def test_step3p5_first_pass_forwards_rejected_token_counts() -> None:
|
||||
# set_inputs_first_pass is inherited from the base proposer; the base
|
||||
# simple-path is the canonical Step3.5 behaviour. Step3.5 only needs to
|
||||
# forward num_rejected_tokens_gpu through _propose.
|
||||
set_inputs = _method(BASE_PROPOSER, "AscendSpecDecodeBaseProposer", "set_inputs_first_pass")
|
||||
propose = _method(STEP3P5, "AscendStep3p5MTPProposer", "_propose")
|
||||
|
||||
assert "num_rejected_tokens_gpu" in [arg.arg for arg in set_inputs.args.args]
|
||||
assert "num_rejected_tokens_gpu=num_rejected_tokens_gpu" in _src(propose)
|
||||
|
||||
# Guard against the override creeping back: the previous step3p5 simple-path
|
||||
# was byte-equivalent to the base's `not needs_extra_input_slots and
|
||||
# pcp_size <= 1` branch, so a re-override is almost certainly redundant.
|
||||
step_methods = {n.name for n in _class(STEP3P5, "AscendStep3p5MTPProposer").body if isinstance(n, ast.FunctionDef)}
|
||||
assert "set_inputs_first_pass" not in step_methods
|
||||
|
||||
|
||||
def test_step3p5_draft_window_and_config_contracts() -> None:
|
||||
base_run = _method(BASE_PROPOSER, "AscendSpecDecodeBaseProposer", "_run_merged_draft")
|
||||
step_run = _method(STEP3P5, "AscendStep3p5MTPProposer", "_run_merged_draft")
|
||||
run_window = _src(_method(STEP3P5, "AscendStep3p5MTPProposer", "_run_window_draft_steps"))
|
||||
build_metadata = _src(_method(STEP3P5, "AscendStep3p5MTPProposer", "_build_step_attn_metadatas"))
|
||||
roll_inputs = _src(_method(STEP3P5, "AscendStep3p5MTPProposer", "_roll_window_inputs_only"))
|
||||
ensure_layer_types = _src(
|
||||
_method(
|
||||
STEP3P5,
|
||||
"AscendStep3p5MTPProposer",
|
||||
"_ensure_draft_layer_types_cover_mtp_layers",
|
||||
)
|
||||
)
|
||||
create_config = _src(_method(STEP3P5, "AscendStep3p5MTPProposer", "_create_draft_vllm_config"))
|
||||
|
||||
assert [arg.arg for arg in step_run.args.args] == [arg.arg for arg in base_run.args.args]
|
||||
assert "multi_steps_attn_metadata.append(per_step_attn_metadata)" in build_metadata
|
||||
assert "multi_steps_attn_metadata[spec_step_idx]" in run_window
|
||||
assert "self.input_ids[token_indices_to_sample]" in roll_inputs
|
||||
assert "_ensure_draft_layer_types_cover_mtp_layers()" in create_config
|
||||
assert "self.draft_model_config.hf_config" in ensure_layer_types
|
||||
assert "self.vllm_config.model_config.hf_config" not in ensure_layer_types
|
||||
assert "sliding_attention" in ensure_layer_types
|
||||
|
||||
|
||||
def test_step3p7_uses_step3p5_mtp_override_without_legacy_runtime_patch() -> None:
|
||||
override_src = _src(_func(PATCH_SPEC_CFG, "hf_config_override"))
|
||||
|
||||
assert "step3p7" in override_src
|
||||
assert "Step3p7ForConditionalGeneration" in override_src
|
||||
assert "step3p5_mtp" in override_src
|
||||
assert "Step3p5MTP" in override_src
|
||||
assert "patch_step3p5_mtp" not in WORKER_PATCH_INIT.read_text()
|
||||
assert not LEGACY_STEP3P7_PATCH.exists()
|
||||
186
tests/ut/spec_decode/test_utils.py
Normal file
186
tests/ut/spec_decode/test_utils.py
Normal file
@@ -0,0 +1,186 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for vllm_ascend.spec_decode.utils.
|
||||
|
||||
These exercise the CPU/GPU correction helpers used by the async spec-decode
|
||||
path on Ascend.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from vllm_ascend.spec_decode.utils import (
|
||||
correct_optimistic_seq_lens_cpu,
|
||||
update_num_computed_tokens_for_batch_change,
|
||||
)
|
||||
|
||||
|
||||
def _build_optimistic_seq_lens(
|
||||
prev_step_computed: np.ndarray,
|
||||
prev_drafts: np.ndarray,
|
||||
num_scheduled_step_n: np.ndarray,
|
||||
) -> np.ndarray:
|
||||
"""Recreate the value the scheduler would write into ``optimistic_seq_lens_cpu``.
|
||||
|
||||
The scheduler advances ``num_computed_tokens_cpu`` by the previous step's
|
||||
scheduled-token count (``prev_drafts + 1``), which is the optimistic count
|
||||
assuming all drafts were accepted.
|
||||
"""
|
||||
optimistic_num_computed = prev_step_computed + (prev_drafts + 1)
|
||||
return optimistic_num_computed + num_scheduled_step_n
|
||||
|
||||
|
||||
def _ground_truth_seq_lens(
|
||||
prev_step_computed: np.ndarray,
|
||||
valid_count: np.ndarray,
|
||||
num_scheduled_step_n: np.ndarray,
|
||||
) -> np.ndarray:
|
||||
"""``self.seq_lens`` (GPU) carries this exact value after correction."""
|
||||
return prev_step_computed + valid_count + num_scheduled_step_n
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_matches_gpu_seq_lens():
|
||||
"""CPU correction must match the seq_lens that the GPU path computes."""
|
||||
num_reqs = 5
|
||||
# prev_step_computed is C[N-2], i.e. num_computed at end of step N-2
|
||||
prev_step_computed = np.array([100, 200, 50, 0, 300], dtype=np.int32)
|
||||
# Drafts scheduled in step N-1; participating reqs have prev_drafts > 0
|
||||
prev_drafts = np.array([2, 4, 0, 0, 3], dtype=np.int32)
|
||||
# Drafts actually accepted; valid_count == 1 + accepted (bonus + accepted)
|
||||
accepted = np.array([2, 0, 0, 0, 1], dtype=np.int32)
|
||||
valid_count = accepted + 1
|
||||
# Step N's scheduled count
|
||||
num_scheduled_step_n = np.array([3, 5, 1, 1, 4], dtype=np.int32)
|
||||
# prev_positions: -1 for new requests, otherwise gather index
|
||||
prev_positions = np.array([0, 1, 2, -1, 4], dtype=np.int32)
|
||||
|
||||
optimistic = _build_optimistic_seq_lens(prev_step_computed, prev_drafts, num_scheduled_step_n).astype(np.int32)
|
||||
|
||||
correct_optimistic_seq_lens_cpu(
|
||||
optimistic,
|
||||
prev_positions,
|
||||
prev_drafts,
|
||||
valid_count.astype(np.int32),
|
||||
num_reqs,
|
||||
)
|
||||
|
||||
expected = _ground_truth_seq_lens(prev_step_computed, valid_count, num_scheduled_step_n)
|
||||
# Non-participating requests (prev_drafts == 0 or prev_positions < 0) keep
|
||||
# the optimistic value, which already coincides with the truth because
|
||||
# there were no drafts to reject.
|
||||
non_participating = (prev_drafts == 0) | (prev_positions < 0)
|
||||
expected[non_participating] = _build_optimistic_seq_lens(
|
||||
prev_step_computed[non_participating],
|
||||
prev_drafts[non_participating],
|
||||
num_scheduled_step_n[non_participating],
|
||||
)
|
||||
np.testing.assert_array_equal(optimistic, expected)
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_no_participants():
|
||||
"""No participating reqs → optimistic_seq_lens unchanged."""
|
||||
optimistic = np.array([10, 20, 30], dtype=np.int32)
|
||||
correct_optimistic_seq_lens_cpu(
|
||||
optimistic,
|
||||
np.array([-1, -1, -1], dtype=np.int32),
|
||||
np.zeros(3, dtype=np.int32),
|
||||
np.zeros(3, dtype=np.int32),
|
||||
3,
|
||||
)
|
||||
np.testing.assert_array_equal(optimistic, np.array([10, 20, 30]))
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_all_accepted():
|
||||
"""When every draft was accepted, the correction is a no-op."""
|
||||
num_reqs = 3
|
||||
prev_drafts = np.array([2, 3, 1], dtype=np.int32)
|
||||
valid_count = (prev_drafts + 1).astype(np.int32) # all accepted
|
||||
prev_positions = np.array([0, 1, 2], dtype=np.int32)
|
||||
optimistic = np.array([105, 208, 51], dtype=np.int32)
|
||||
expected = optimistic.copy()
|
||||
correct_optimistic_seq_lens_cpu(optimistic, prev_positions, prev_drafts, valid_count, num_reqs)
|
||||
np.testing.assert_array_equal(optimistic, expected)
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_all_rejected():
|
||||
"""When every draft was rejected (only bonus token kept), correction == prev_drafts."""
|
||||
num_reqs = 3
|
||||
prev_drafts = np.array([2, 3, 1], dtype=np.int32)
|
||||
valid_count = np.array([1, 1, 1], dtype=np.int32) # only bonus kept
|
||||
prev_positions = np.array([0, 1, 2], dtype=np.int32)
|
||||
optimistic = np.array([105, 208, 51], dtype=np.int32)
|
||||
expected = optimistic - prev_drafts
|
||||
correct_optimistic_seq_lens_cpu(optimistic, prev_positions, prev_drafts, valid_count, num_reqs)
|
||||
np.testing.assert_array_equal(optimistic, expected)
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_in_place_mutation():
|
||||
"""The function must modify the input array in place."""
|
||||
optimistic = np.array([100, 200], dtype=np.int64)
|
||||
prev_drafts = np.array([2, 0], dtype=np.int32)
|
||||
valid_count = np.array([2, 0], dtype=np.int32) # 1 of 2 accepted, prefill
|
||||
prev_positions = np.array([0, -1], dtype=np.int32)
|
||||
pre_id = id(optimistic)
|
||||
correct_optimistic_seq_lens_cpu(optimistic, prev_positions, prev_drafts, valid_count, 2)
|
||||
assert id(optimistic) == pre_id
|
||||
# req 0: correction = prev_drafts + 1 - valid_count = 2 + 1 - 2 = 1
|
||||
# req 1: not participating → unchanged
|
||||
np.testing.assert_array_equal(optimistic, np.array([99, 200]))
|
||||
|
||||
|
||||
def test_correct_optimistic_seq_lens_cpu_partial_batch():
|
||||
"""Only the first num_reqs entries of the buffer are touched."""
|
||||
optimistic = np.array([100, 200, 999, 999], dtype=np.int32)
|
||||
prev_drafts = np.array([2, 1, 0, 0], dtype=np.int32)
|
||||
valid_count = np.array([2, 2, 0, 0], dtype=np.int32)
|
||||
prev_positions = np.array([0, 1, -1, -1], dtype=np.int32)
|
||||
|
||||
correct_optimistic_seq_lens_cpu(optimistic, prev_positions, prev_drafts, valid_count, 2)
|
||||
# req 0: correction = 2 + 1 - 2 = 1, → 99
|
||||
# req 1: correction = 1 + 1 - 2 = 0, → 200 unchanged
|
||||
# tail (idx 2,3) untouched
|
||||
np.testing.assert_array_equal(optimistic, np.array([99, 200, 999, 999]))
|
||||
|
||||
|
||||
def test_cpu_and_gpu_corrections_agree():
|
||||
"""The CPU helper must agree with ``update_num_computed_tokens_for_batch_change``.
|
||||
|
||||
They live on different sides of the device boundary, but both implement
|
||||
the same correction. We compare the post-correction seq_lens (= corrected
|
||||
num_computed_tokens + num_scheduled_step_n) on each path.
|
||||
"""
|
||||
num_reqs = 6
|
||||
prev_step_computed = np.array([100, 250, 80, 0, 410, 5], dtype=np.int32)
|
||||
prev_drafts = np.array([2, 4, 0, 0, 3, 1], dtype=np.int32)
|
||||
accepted = np.array([2, 1, 0, 0, 0, 1], dtype=np.int32)
|
||||
valid_count = (accepted + 1).astype(np.int32)
|
||||
num_scheduled_step_n = np.array([3, 5, 1, 1, 4, 2], dtype=np.int32)
|
||||
prev_positions = np.array([0, 1, 2, -1, 4, 5], dtype=np.int32)
|
||||
|
||||
# CPU path
|
||||
optimistic = _build_optimistic_seq_lens(prev_step_computed, prev_drafts, num_scheduled_step_n).astype(np.int32)
|
||||
correct_optimistic_seq_lens_cpu(optimistic, prev_positions, prev_drafts, valid_count, num_reqs)
|
||||
|
||||
# GPU path on CPU device for portability
|
||||
cpu_num_computed = torch.from_numpy(
|
||||
prev_step_computed + prev_drafts + 1 # scheduler-bumped optimistic
|
||||
).to(torch.int32)
|
||||
num_computed_gpu = torch.from_numpy(prev_step_computed.copy()).to(torch.int32)
|
||||
num_accepted_gpu = torch.zeros(num_reqs, dtype=torch.int32)
|
||||
valid_count_t = torch.from_numpy(valid_count)
|
||||
prev_positions_t = torch.from_numpy(prev_positions)
|
||||
prev_drafts_t = torch.from_numpy(prev_drafts)
|
||||
update_num_computed_tokens_for_batch_change(
|
||||
num_computed_gpu,
|
||||
num_accepted_gpu,
|
||||
prev_positions_t,
|
||||
valid_count_t,
|
||||
prev_drafts_t,
|
||||
cpu_num_computed,
|
||||
)
|
||||
|
||||
gpu_seq_lens = num_computed_gpu.numpy() + num_scheduled_step_n
|
||||
|
||||
np.testing.assert_array_equal(optimistic, gpu_seq_lens)
|
||||
Reference in New Issue
Block a user