init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

View File

File diff suppressed because it is too large Load Diff

View 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

View 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()

View 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)

View 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()

View 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)