Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions tests/ut/worker/test_attn_utils_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -647,6 +647,7 @@ def test_dsv4_backends_declare_role_specific_logical_sizes(
("model_state", CUDAGraphMode.NONE, False, 1, 5),
("model_state", CUDAGraphMode.FULL, False, 1, 8),
("pcp_capture", CUDAGraphMode.NONE, True, 2, 8),
("pcp_runtime", CUDAGraphMode.NONE, False, 2, 8),
],
)
def test_mrv2_builds_shared_dsa_metadata_for_each_execution_mode(
Expand All @@ -666,7 +667,7 @@ def test_mrv2_builds_shared_dsa_metadata_for_each_execution_mode(
[2, 1, 0, 0],
dtype=torch.int32,
)
pcp_context = object() if caller == "pcp_capture" else None
pcp_context = object() if pcp_size > 1 else None
pcp_manager = (
SimpleNamespace(
build_attention_context=MagicMock(return_value=pcp_context),
Expand Down Expand Up @@ -748,7 +749,7 @@ def test_mrv2_builds_shared_dsa_metadata_for_each_execution_mode(
if pcp_context is not None:
assert [call["pcp_cache_group_idx"] for call in calls] == [0, 1]
assert pcp_manager is not None
pcp_manager.build_attention_context.assert_called_once_with(input_batch)
pcp_manager.build_attention_context.assert_called_once_with(input_batch, block_tables, slot_mappings)
else:
assert all(call["pcp_cache_group_idx"] is None for call in calls)

Expand Down
42 changes: 42 additions & 0 deletions tests/ut/worker/test_model_runner_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
from vllm.config import CUDAGraphMode
from vllm.v1.worker.gpu.model_runner import GPUModelRunner

from vllm_ascend.worker.v2.input_batch import AscendInputBatch, AscendInputBuffers
from vllm_ascend.worker.v2.model_runner import NPUModelRunner
from vllm_ascend.worker.v2.pcp_manager import AscendPCPManager


def _make_runner(need_timing: bool = True):
Expand Down Expand Up @@ -214,3 +216,43 @@ def test_prepare_inputs_preserves_pcp_tokens_and_forwards_graph_padding():
assert padded_num_tokens.attr == "num_tokens"
assert isinstance(padded_num_tokens.value, ast.Name)
assert padded_num_tokens.value.id == "batch_desc"


@pytest.mark.parametrize("num_reqs,num_tokens", [(4, 4), (2, 6)])
def test_pcp_dummy_refreshes_captured_buffers_after_real_batch(num_reqs, num_tokens):
runner = _make_runner()
runner.input_buffers = AscendInputBuffers(4, 8, torch.device("cpu"))
manager = AscendPCPManager(2, 1, torch.device("cpu"), max_num_reqs=4, max_num_tokens=8)
runner.pcp_manager = manager
manager._local_block_tables = (torch.full((8, 2), 99, dtype=torch.int32),)
manager._gathered_kv_slot_mappings = torch.full((1, 16), 99, dtype=torch.int64)
captured = {
name: getattr(manager.input_buffers, name)
for name in ("input_ids", "positions", "is_padding", "query_start_loc", "seq_lens")
}
for name, value in captured.items():
value.fill_(False if name == "is_padding" else 99)
manager.input_buffers.seq_lens_np.fill(99)
with patch("vllm_ascend.worker.v2.input_batch.update_cos_sin"):
dummy = AscendInputBatch.make_dummy(num_reqs, num_tokens, runner.input_buffers)

block_tables, slots = runner.prepare_dummy_attn(dummy)

for name, value in captured.items():
expected = getattr(dummy, name)
torch.testing.assert_close(value[: len(expected)], expected)
np.testing.assert_array_equal(manager.input_buffers.seq_lens_np[:num_reqs], dummy.seq_lens_np)
assert block_tables[0].data_ptr() == manager._local_block_tables[0].data_ptr()
assert torch.count_nonzero(block_tables[0]) == 0
assert slots.data_ptr() == manager._gathered_kv_slot_mappings.data_ptr()
assert slots.shape == (1, 2 * num_tokens)
assert torch.all(slots == -1)


def test_prepare_dummy_attn_without_pcp_uses_upstream():
runner = _make_runner()
runner.pcp_manager = None
dummy = object()
with patch.object(GPUModelRunner, "prepare_dummy_attn", return_value=((), None)) as parent:
assert runner.prepare_dummy_attn(dummy) == ((), None)
parent.assert_called_once_with(dummy)
67 changes: 62 additions & 5 deletions tests/ut/worker/test_mtp_pcp_speculator_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,16 @@
import numpy as np
import pytest
import torch
from vllm.config import CUDAGraphMode
from vllm.v1.worker.gpu import dp_utils
from vllm.v1.worker.gpu.spec_decode.eagle.speculator import EagleSpeculator
from vllm.v1.worker.gpu.spec_decode.mtp.speculator import MTPSpeculator

from vllm_ascend.worker.v2.input_batch import AscendInputBatch
from vllm_ascend.worker.v2.spec_decode.autoregressive import (
speculator as speculator_module,
)
from vllm_ascend.worker.v2.spec_decode.eagle.speculator import AscendEagleSpeculator
from vllm_ascend.worker.v2.spec_decode.mtp.speculator import (
AscendMTPSpeculator,
)
Expand Down Expand Up @@ -322,25 +326,61 @@ def test_graph_prefill_builds_draft_metadata(replicated_pcp: bool) -> None:
)


def test_propose_disables_target_pcp_manager_for_replicated_draft() -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = True
@pytest.mark.skipif(speculator_module.vllm_version_is("0.28.0"), reason="DPSyncState is a main2main interface")
@pytest.mark.parametrize(
("speculator_cls", "parent_cls", "replicated_pcp", "batch_kind"),
[
(AscendMTPSpeculator, MTPSpeculator, True, "prefill"),
(AscendMTPSpeculator, MTPSpeculator, True, "decode"),
(AscendMTPSpeculator, MTPSpeculator, True, "idle"),
(AscendMTPSpeculator, MTPSpeculator, False, "prefill"),
(AscendEagleSpeculator, EagleSpeculator, True, "prefill"),
],
)
def test_propose_sync_follows_draft_token_layout(speculator_cls, parent_cls, replicated_pcp, batch_kind) -> None:
speculator = object.__new__(speculator_cls)
speculator.replicated_pcp = replicated_pcp
speculator.input_batch = None
speculator.pcp_manager = MagicMock()
speculator.model_state = SimpleNamespace(
pcp_manager=speculator.pcp_manager,
)
input_batch = _make_padded_input_batch()
input_batch.has_prefill = batch_kind == "prefill"
input_batch.is_dummy = batch_kind == "idle"
if batch_kind != "prefill":
input_batch.num_tokens_after_padding = 2
num_tokens = input_batch.num_tokens_after_padding
target_num_tokens = num_tokens // 2 if replicated_pcp and input_batch.has_prefill else num_tokens
target_sync = SimpleNamespace(
eager=True,
uniform_token_count=None,
num_tokens_across_dp=torch.tensor([target_num_tokens, target_num_tokens]),
)
expected = object()

def parent_propose(*args, **kwargs):
assert args[0] is input_batch
assert speculator.model_state.pcp_manager is None
assert (speculator.model_state.pcp_manager is None) is replicated_pcp
# Exercise upstream's real reuse checks; only the collective is mocked.
with patch.object(dp_utils, "sync_cudagraph_and_dp_padding") as sync:
sync.return_value = (SimpleNamespace(cg_mode=CUDAGraphMode.NONE), object())
dp_utils.dispatch_cg_and_sync_dp(
None,
input_batch.num_reqs,
num_tokens,
None,
dp_size=2,
dp_rank=0,
need_eager=True,
dp_sync=args[11],
)
assert sync.call_count == int(replicated_pcp)
return expected

with (
patch.object(
MTPSpeculator,
parent_cls,
"propose",
side_effect=parent_propose,
),
Expand All @@ -358,8 +398,25 @@ def parent_propose(*args, **kwargs):
actual = speculator.propose(
input_batch,
*[MagicMock() for _ in range(10)],
dp_sync=target_sync,
)

assert actual is expected
assert speculator.input_batch is input_batch
assert speculator.model_state.pcp_manager is speculator.pcp_manager


def test_propose_preserves_v028_dp_token_counts() -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = True
input_batch = object()
token_counts = torch.tensor([4, 8])
with (
patch.object(speculator_module, "vllm_version_is", return_value=True),
patch.object(speculator_module, "disable_target_pcp_for_replicated_draft", return_value=nullcontext()),
patch.object(speculator_module, "build_attn_metadata_wrapper", return_value=nullcontext()),
patch.object(speculator_module, "torch_gather_wrapper", return_value=nullcontext()),
patch.object(MTPSpeculator, "propose") as parent,
):
speculator.propose(input_batch, *[MagicMock() for _ in range(10)], token_counts, dp_sync=object())
assert parent.call_args.args[11] is token_counts
49 changes: 49 additions & 0 deletions tests/ut/worker/test_pcp_manager_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,12 +61,14 @@ def _make_pcp_config(
*,
sparse_mla: bool = True,
pipeline_parallel_size: int = 1,
data_parallel_size: int = 1,
):
hf_text_config = SimpleNamespace(index_topk=2048) if sparse_mla else SimpleNamespace()
return SimpleNamespace(
parallel_config=SimpleNamespace(
prefill_context_parallel_size=2,
pipeline_parallel_size=pipeline_parallel_size,
data_parallel_size=data_parallel_size,
),
model_config=SimpleNamespace(
use_mla=True,
Expand Down Expand Up @@ -980,3 +982,50 @@ def test_partition_batch_clears_padded_dcp_local_seq_lens() -> None:
result.dcp_local_seq_lens,
torch.tensor([4, 5, 0, 0, 0, 0, 0, 0], dtype=torch.int32),
)


@pytest.mark.parametrize(
"dp_size,cudagraph_mode,allowed",
[
(2, CUDAGraphMode.NONE, True),
(2, CUDAGraphMode.FULL_DECODE_ONLY, True),
(2, CUDAGraphMode.PIECEWISE, False),
(1, CUDAGraphMode.PIECEWISE, True),
],
)
def test_validate_config_pcp_dp_graph_modes(dp_size, cudagraph_mode, allowed):
config = _make_pcp_config(cudagraph_mode, sparse_mla=False, data_parallel_size=dp_size)
if allowed:
AscendPCPManager.validate_config(config, supports_mm_inputs=False)
else:
with pytest.raises(NotImplementedError, match=r"PCP\+DP supports eager mode or FULL_DECODE_ONLY"):
AscendPCPManager.validate_config(config, supports_mm_inputs=False)


@pytest.mark.parametrize("pcp_rank", [0, 1])
@pytest.mark.parametrize("has_stale_batch", [False, True])
def test_dummy_attention_context_uses_current_batch(pcp_rank, has_stale_batch):
manager = AscendPCPManager(2, pcp_rank, torch.device("cpu"))
saved_batch = _make_global_pcp_batch() if has_stale_batch else None
manager._global_batch = saved_batch
manager._hidden_restore_idx = torch.tensor([99]) if has_stale_batch else None
saved_indices = manager._hidden_restore_idx
manager._block_tables = SimpleNamespace(gather_block_tables=MagicMock())
manager._global_batch_slot_mappings = torch.full((2, 32), 99, dtype=torch.int64)
dummy = _make_local_pcp_batch()
dummy.is_dummy = True
dummy.num_tokens = 4 # Exercise padding: the layout stride must still be 6.
block_tables = (torch.zeros((2, 1), dtype=torch.int32),) * 2
slot_mappings = torch.arange(24, dtype=torch.int64).reshape(2, 12)

context = manager.build_attention_context(dummy, block_tables, slot_mappings)

assert context.global_batch is dummy
assert context.global_block_tables is block_tables
start = pcp_rank * 6
torch.testing.assert_close(context.global_slot_mappings, slot_mappings[:, start : start + 6])
gathered_hidden = torch.arange(12).reshape(12, 1)
torch.testing.assert_close(gathered_hidden[context.hidden_restore_idx], gathered_hidden[start : start + 6])
assert manager._global_batch is saved_batch
assert manager._hidden_restore_idx is saved_indices
manager._block_tables.gather_block_tables.assert_not_called()
5 changes: 5 additions & 0 deletions vllm_ascend/worker/v2/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -534,6 +534,11 @@ def prepare_inputs( # type: ignore[misc]

return input_batch

def prepare_dummy_attn(self, input_batch: AscendInputBatch) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
if self.pcp_manager is None:
return super().prepare_dummy_attn(input_batch)
return self.pcp_manager.prepare_dummy_attn(input_batch)

def _lmhead_tp_max_num_logits(self) -> int:
"""Logits row capacity shared by every rank of the lmhead-TP group.

Expand Down
6 changes: 5 additions & 1 deletion vllm_ascend/worker/v2/model_states/default.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,11 @@ def prepare_attn(
query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np)
is_prefilling = torch.from_numpy(input_batch.is_prefilling_np)
max_query_len = input_batch.num_scheduled_tokens.max().item()
pcp_context = self.pcp_manager.build_attention_context(input_batch) if self.pcp_manager is not None else None
pcp_context = (
self.pcp_manager.build_attention_context(input_batch, block_tables, slot_mappings)
if self.pcp_manager is not None
else None
)
# attn_metadata is needed when update_full_graph_params, but no way can get it now.
# Temporarily store it in model_state.
self.attn_metadata = build_attn_metadata(
Expand Down
68 changes: 46 additions & 22 deletions vllm_ascend/worker/v2/pcp_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,11 @@ def validate_config(
)
is_sparse_mla = hasattr(model_config.hf_text_config, "index_topk")
cudagraph_mode = vllm_config.compilation_config.cudagraph_mode
if parallel_config.data_parallel_size > 1 and cudagraph_mode not in {
CUDAGraphMode.NONE,
CUDAGraphMode.FULL_DECODE_ONLY,
}:
raise NotImplementedError("MRV2 PCP+DP supports eager mode or FULL_DECODE_ONLY CUDA graphs only.")
if is_sparse_mla and cudagraph_mode not in {
CUDAGraphMode.NONE,
CUDAGraphMode.FULL_DECODE_ONLY,
Expand Down Expand Up @@ -354,6 +359,26 @@ def restore_hidden_state_buffer(self, hidden_states: torch.Tensor) -> None:
restored_hidden_states = self.restore_hidden_states(hidden_states[:local_num_tokens_padded])
hidden_states[: restored_hidden_states.shape[0]].copy_(restored_hidden_states)

# TODO(wzx0726): Once the paired vLLM includes https://github.com/vllm-project/vllm/pull/53867,
# adapt its PCP prepare_inputs_to_capture path to create AscendInputBatch
# directly in persistent PCP buffers, then remove this method and the
# NPUModelRunner.prepare_dummy_attn override after capture/idle replay validation.
def prepare_dummy_attn(self, input_batch: AscendInputBatch) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
# Runtime dummy inputs use the runner buffers, whereas FULL graphs
# capture PCP-local storage. Refresh that storage after a real batch.
input_buffers = self.input_buffers
num_tokens = input_batch.num_tokens_after_padding
num_reqs = input_batch.num_reqs_after_padding
for name in ("input_ids", "positions", "is_padding"):
getattr(input_buffers, name)[:num_tokens].copy_(getattr(input_batch, name))
input_buffers.query_start_loc[: num_reqs + 1].copy_(input_batch.query_start_loc)
input_buffers.seq_lens[:num_reqs].copy_(input_batch.seq_lens)
input_buffers.seq_lens_np[:num_reqs] = input_batch.seq_lens_np[:num_reqs]
return (
self.get_dummy_block_tables(num_reqs),
self.get_dummy_slot_mappings(num_tokens),
)

def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]:
"""Return capture views backed by the persistent PCP-local tables.

Expand Down Expand Up @@ -415,34 +440,33 @@ def prepare_slot_mappings(self) -> torch.Tensor:

def build_attention_context(
self,
capture_batch: AscendInputBatch | None = None,
input_batch: AscendInputBatch | None = None,
block_tables: tuple[torch.Tensor, ...] | None = None,
slot_mappings: torch.Tensor | None = None,
) -> AscendPCPAttentionContext:
"""Build the PCP context consumed by attention metadata builders.
"""Build PCP context for the current real, capture, or idle DP batch."""
if input_batch is not None and input_batch.is_dummy:
# Both capture and runtime dummy batches bypass partition_batch().
# Saved layout state may be absent or belong to a previous request.
assert block_tables is not None
assert slot_mappings is not None
num_tokens = input_batch.num_tokens_after_padding
restore_start = self.pcp_rank * num_tokens
return AscendPCPAttentionContext(
global_batch=input_batch,
global_block_tables=block_tables,
global_slot_mappings=slot_mappings.view(slot_mappings.shape[0], self.pcp_world_size, num_tokens)[
:, self.pcp_rank
],
Comment on lines +458 to +460

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

Using .view() on slot_mappings can raise a RuntimeError if the tensor is non-contiguous (e.g., if it was created via slicing or other non-contiguous operations). Since slot_mappings is passed from external callers and its contiguity is not guaranteed, it is safer to use .reshape() instead of .view().

Suggested change
global_slot_mappings=slot_mappings.view(slot_mappings.shape[0], self.pcp_world_size, num_tokens)[
:, self.pcp_rank
],
global_slot_mappings=slot_mappings.reshape(slot_mappings.shape[0], self.pcp_world_size, num_tokens)[
:, self.pcp_rank
],

hidden_restore_idx=torch.arange(restore_start, restore_start + num_tokens, device=self.device),
)

At runtime the global batch partitioned by ``partition_batch`` is the
authoritative view. Graph capture (vLLM #53869) never partitions a
batch, so callers pass the capture-only dummy batch laid out on the
same persistent input buffers.
"""
global_batch = self._global_batch
if global_batch is None:
global_batch = capture_batch
assert global_batch is not None
# Only duck-typed attributes are consumed downstream, so callers may
# pass batch stand-ins (e.g. UT SimpleNamespace or the capture dummy).
hidden_restore_idx = self._hidden_restore_idx
if hidden_restore_idx is None:
# Graph capture (vLLM #53515/#53869) never partitions a batch, so
# _build_batch_layout did not fill _hidden_restore_idx. Runtime
# prepare_attn rebuilds this metadata every step, so an identity
# placeholder is sufficient for the capture-only context.
hidden_restore_idx = torch.arange(
global_batch.num_tokens_after_padding,
dtype=torch.int64,
device=self.device,
)
assert global_batch is not None
assert self._block_tables is not None
assert self._global_batch_slot_mappings is not None
assert hidden_restore_idx is not None
return AscendPCPAttentionContext(
global_batch=global_batch,
global_block_tables=self._block_tables.gather_block_tables(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -255,7 +255,12 @@ def propose(
generate_draft.
"""
self.input_batch = input_batch
sync_state = num_tokens_across_dp if vllm_version_is("0.28.0") else dp_sync
if vllm_version_is("0.28.0"):
sync_state = num_tokens_across_dp
else:
# Replicated drafts use global tokens, unlike the PCP-local target.
# Every DP rank must take the draft sync, including decode and idle ranks.
sync_state = None if self.replicated_pcp else dp_sync
# wrap build_attn_metadata to use Ascend attention metadata building.
# so we can call super().propose() directly.
with (
Expand Down
Loading