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
145 changes: 130 additions & 15 deletions tests/ut/worker/test_mtp_pcp_speculator_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,8 +195,11 @@ def test_draft_prefill_attn_groups_follow_draft_topology(
assert speculator.draft_prefill_attn_groups is expected


def test_prepare_replicated_prefill_attn_uses_global_batch() -> None:
@pytest.mark.parametrize("attn_architecture", ["GQA", "MLA", "DSA", "SFA"])
@pytest.mark.parametrize("cudagraph_runtime_mode", [CUDAGraphMode.NONE, CUDAGraphMode.PIECEWISE])
def test_prepare_replicated_prefill_attn_uses_global_batch(attn_architecture, cudagraph_runtime_mode) -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.attn_architecture = attn_architecture
speculator.block_tables = MagicMock()
speculator.kv_cache_config = object()
speculator._build_draft_attn_metadata = MagicMock(return_value={"draft.layer": object()})
Expand All @@ -220,6 +223,7 @@ def test_prepare_replicated_prefill_attn_uses_global_batch() -> None:
original_slot_mappings,
input_batch.num_reqs_after_padding,
input_batch.num_tokens_after_padding,
cudagraph_runtime_mode=cudagraph_runtime_mode,
)

assert attn_metadata == speculator._build_draft_attn_metadata.return_value
Expand Down Expand Up @@ -248,7 +252,8 @@ def test_prepare_replicated_prefill_attn_uses_global_batch() -> None:
)


def test_prefill_rebuilds_replicated_pcp_metadata_before_filtering() -> None:
@pytest.mark.parametrize("cudagraph_runtime_mode", [CUDAGraphMode.NONE, CUDAGraphMode.PIECEWISE])
def test_prefill_rebuilds_replicated_pcp_metadata_before_filtering(cudagraph_runtime_mode) -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = True
speculator.input_batch = _make_padded_input_batch()
Expand Down Expand Up @@ -279,13 +284,15 @@ def test_prefill_rebuilds_replicated_pcp_metadata_before_filtering() -> None:
attn_metadata=local_attn_metadata,
slot_mappings=local_slot_mappings,
num_tokens_across_dp=None,
cudagraph_runtime_mode=cudagraph_runtime_mode,
)

speculator._prepare_replicated_prefill_attn.assert_called_once_with(
local_attn_metadata,
local_slot_mappings,
2,
8,
cudagraph_runtime_mode=cudagraph_runtime_mode,
)
parent_prefill.assert_called_once()
parent_args = parent_prefill.call_args.args
Expand All @@ -294,18 +301,9 @@ def test_prefill_rebuilds_replicated_pcp_metadata_before_filtering() -> None:
assert parent_args[3] is global_slot_mappings


@pytest.mark.parametrize(
("attn_architecture", "replicated_pcp", "rebuild_metadata"),
[
("DSA", True, True),
("DSA", False, False),
("SFA", True, True),
("MLA", True, False),
],
)
def test_graph_prefill_builds_draft_metadata(
attn_architecture: str, replicated_pcp: bool, rebuild_metadata: bool
) -> None:
@pytest.mark.parametrize("attn_architecture", ["GQA", "MLA", "DSA", "SFA"])
@pytest.mark.parametrize("replicated_pcp", [False, True])
def test_graph_prefill_builds_draft_metadata(attn_architecture: str, replicated_pcp: bool) -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = replicated_pcp
speculator.attn_architecture = attn_architecture
Expand All @@ -323,19 +321,136 @@ def test_graph_prefill_builds_draft_metadata(
return_value={"draft.layer": global_draft_metadata},
)

with patch.object(speculator_module, "build_slot_mappings_by_layer", return_value={}):
with patch.object(speculator_module, "build_slot_mappings_by_layer", return_value={}) as build_slots:
actual = speculator.build_draft_attn_metadatas(
num_reqs_padded=2,
num_tokens_padded=8,
is_draft_model_prefill=True,
)

rebuild_metadata = replicated_pcp and attn_architecture in ("DSA", "SFA")
expected_metadata = global_draft_metadata if rebuild_metadata else local_draft_metadata
assert actual == [{"draft.layer": expected_metadata}]
assert actual[0]["draft.layer"] is expected_metadata
assert speculator._build_draft_attn_metadata.call_count == int(rebuild_metadata)
assert build_slots.call_count == int(rebuild_metadata)
assert speculator.block_tables.gather_block_tables.call_count == int(replicated_pcp)
assert speculator.block_tables.compute_slot_mappings.call_count == int(replicated_pcp)
assert local_draft_metadata.decode.actual_seq_lengths_q[-1] == 8


@pytest.mark.parametrize("attn_architecture", ["GQA", "MLA"])
def test_graph_prefill_refreshes_captured_cache_buffers(attn_architecture: str) -> None:
speculator = object.__new__(AscendEagleSpeculator)
speculator.replicated_pcp = True
speculator.attn_architecture = attn_architecture
speculator.draft_attn_layer_names = {"draft.layer"}
metadata = SimpleNamespace(actual_seq_lengths_q=[4, 8])
speculator.model_state = SimpleNamespace(attn_metadata={"draft.layer": metadata})
speculator.kv_cache_config = object()
speculator._build_draft_attn_metadata = MagicMock()

# These views stand in for the persistent buffers bound during capture.
captured_blocks = torch.zeros((2, 3), dtype=torch.int32)
captured_slots = torch.full((1, 8), -1, dtype=torch.int32)
block_ptr, slot_ptr = captured_blocks.data_ptr(), captured_slots.data_ptr()
request_blocks = {3: [7, 13, 0], 7: [17, 19, 23]}

def gather_blocks(idx_mapping, num_reqs_padded):
captured_blocks.zero_()
for row, req_idx in enumerate(idx_mapping.tolist()):
captured_blocks[row] = torch.tensor(request_blocks[req_idx])
return (captured_blocks[:num_reqs_padded],)

def compute_slots(idx_mapping, query_start_loc, positions, num_tokens_padded):
captured_slots.fill_(-1)
for row, req_idx in enumerate(idx_mapping.tolist()):
start, end = query_start_loc[row : row + 2].tolist()
for token in range(start, end):
position = int(positions[token])
block = request_blocks[req_idx][position // 128]
captured_slots[0, token] = block * 128 + position % 128
return captured_slots[:, :num_tokens_padded]

speculator.block_tables = SimpleNamespace(
gather_block_tables=gather_blocks,
compute_slot_mappings=compute_slots,
)
for req_idx, positions, expected_slots in [
(3, [126, 127, 128, 129], [1022, 1023, 1664, 1665]),
(7, [254, 255, 256, 257], [2558, 2559, 2944, 2945]),
]:
batch = _make_padded_input_batch()
batch.is_dummy = False
batch.num_reqs, batch.num_tokens = 1, 4
batch.query_start_loc_np = np.array([0, 4, 8], dtype=np.int32)
batch.idx_mapping = torch.tensor([req_idx], dtype=torch.int32)
batch.query_start_loc = torch.tensor([0, 4], dtype=torch.int32)
batch.positions = torch.tensor(positions, dtype=torch.int64)
speculator.input_batch = batch
# The preceding single-token decode has left a different slot layout.
captured_slots.fill_(-1)
captured_slots[0, 0] = 123
captured_blocks.fill_(-1)

with patch.object(speculator_module, "build_slot_mappings_by_layer") as build_slots:
result = speculator.build_draft_attn_metadatas(2, 8, is_draft_model_prefill=True)

assert result[0]["draft.layer"] is metadata
speculator._build_draft_attn_metadata.assert_not_called()
build_slots.assert_not_called()
assert metadata.actual_seq_lengths_q == [4, 8]
assert captured_slots.tolist() == [expected_slots + [-1] * 4]
assert captured_blocks.tolist() == [request_blocks[req_idx], [0, 0, 0]]
assert (captured_blocks.data_ptr(), captured_slots.data_ptr()) == (block_ptr, slot_ptr)


@pytest.mark.parametrize("guard", ["non_pcp", "no_batch", "dummy", "no_metadata"])
def test_prepare_replicated_prefill_preserves_bypass(guard: str) -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = guard != "non_pcp"
speculator.input_batch = _make_padded_input_batch() if guard != "no_batch" else None
if speculator.input_batch is not None:
speculator.input_batch.is_dummy = guard == "dummy"
speculator.block_tables = MagicMock()
speculator._build_draft_attn_metadata = MagicMock()
metadata = None if guard == "no_metadata" else {"draft.layer": object()}
slots = {"draft.layer": object()}

actual_metadata, actual_slots = speculator._prepare_replicated_prefill_attn(
metadata, slots, 2, 8, cudagraph_runtime_mode=CUDAGraphMode.FULL
)

assert actual_metadata is metadata
assert actual_slots is slots
speculator.block_tables.gather_block_tables.assert_not_called()
speculator.block_tables.compute_slot_mappings.assert_not_called()
speculator._build_draft_attn_metadata.assert_not_called()


@pytest.mark.parametrize("attn_architecture", ["GQA", "MLA", "DSA", "SFA"])
@pytest.mark.parametrize("guard", ["no_batch", "dummy"])
def test_graph_prefill_without_real_batch_preserves_metadata(attn_architecture: str, guard: str) -> None:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = True
speculator.attn_architecture = attn_architecture
speculator.input_batch = _make_padded_input_batch() if guard == "dummy" else None
if speculator.input_batch is not None:
speculator.input_batch.is_dummy = True
speculator.block_tables = MagicMock()
speculator._build_draft_attn_metadata = MagicMock()
metadata = object()
speculator.draft_attn_layer_names = {"draft.layer"}
speculator.model_state = SimpleNamespace(attn_metadata={"draft.layer": metadata})

[actual] = speculator.build_draft_attn_metadatas(2, 8, is_draft_model_prefill=True)

assert actual["draft.layer"] is metadata
speculator.block_tables.gather_block_tables.assert_not_called()
speculator.block_tables.compute_slot_mappings.assert_not_called()
speculator._build_draft_attn_metadata.assert_not_called()


@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"),
Expand Down
32 changes: 20 additions & 12 deletions vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,19 +157,21 @@ def _prepare_replicated_prefill_attn(
slot_mappings: dict[str, torch.Tensor] | None,
num_reqs_padded: int,
num_tokens_padded: int,
cudagraph_runtime_mode: CUDAGraphMode,
) -> tuple[
dict[str, Any] | None,
dict[str, torch.Tensor] | None,
]:
"""Rebuild global draft prefill state for replicated PCP."""
"""Refresh global draft mappings and prepare attention for replicated PCP."""
input_batch = self.input_batch
if not self.replicated_pcp or attn_metadata is None or input_batch is None:
if attn_metadata is None or not self.replicated_pcp or input_batch is None:
return attn_metadata, slot_mappings

assert isinstance(input_batch, AscendInputBatch)
if input_batch.is_dummy:
return attn_metadata, slot_mappings

# Omitting out updates the default buffers bound by draft graph capture.
self.block_tables.gather_block_tables(
input_batch.idx_mapping,
num_reqs_padded=num_reqs_padded,
Expand All @@ -180,6 +182,12 @@ def _prepare_replicated_prefill_attn(
input_batch.positions,
num_tokens_padded=num_tokens_padded,
)
# TODO: Remove this early return once FIA supports padded Query tensors
# whose token count exceeds the cumulative query length. Keep the
# mapping refresh above when unifying metadata construction.
if cudagraph_runtime_mode == CUDAGraphMode.FULL and self.attn_architecture in ("MLA", "GQA"):
return attn_metadata, slot_mappings

slot_mappings = build_slot_mappings_by_layer(
slot_mappings_tensor,
self.kv_cache_config,
Expand Down Expand Up @@ -432,6 +440,7 @@ def _prefill(
slot_mappings,
num_reqs,
num_tokens,
cudagraph_runtime_mode=cudagraph_runtime_mode,
)
# Draft prefill reuses target metadata, but the target metadata may
# also contain target-only attention layers (e.g. GDN layers).
Expand Down Expand Up @@ -498,16 +507,15 @@ def build_draft_attn_metadatas(
}

if is_draft_model_prefill:
if self.attn_architecture in ("DSA", "SFA"):
prepared_attn_metadata, _ = self._prepare_replicated_prefill_attn(
attn_metadata,
None,
num_reqs_padded,
num_tokens_padded,
)
assert prepared_attn_metadata is not None
attn_metadata = prepared_attn_metadata
return [attn_metadata]
prepared_attn_metadata, _ = self._prepare_replicated_prefill_attn(
attn_metadata,
None,
num_reqs_padded,
num_tokens_padded,
cudagraph_runtime_mode=CUDAGraphMode.FULL,
)
assert prepared_attn_metadata is not None
return [prepared_attn_metadata]

draft_attn_metadatas = self._init_decode_draft_attn_metadatas(attn_metadata, num_reqs_padded)

Expand Down
Loading