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
56 changes: 31 additions & 25 deletions tests/ut/worker/test_mtp_pcp_speculator_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,40 +294,46 @@ def test_prefill_rebuilds_replicated_pcp_metadata_before_filtering() -> None:
assert parent_args[3] is global_slot_mappings


@pytest.mark.parametrize("replicated_pcp", [False, True])
def test_graph_prefill_builds_draft_metadata(replicated_pcp: bool) -> None:
@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:
speculator = object.__new__(AscendMTPSpeculator)
speculator.replicated_pcp = replicated_pcp
speculator.attn_architecture = attn_architecture
speculator.input_batch = _make_padded_input_batch()
speculator.input_batch.is_dummy = False
speculator.block_tables = MagicMock()
speculator.kv_cache_config = object()
speculator.draft_attn_layer_names = {"draft.layer"}
local_draft_metadata = object()
local_draft_metadata = SimpleNamespace(decode=SimpleNamespace(actual_seq_lengths_q=[4, 8]))
global_draft_metadata = object()
speculator.model_state = SimpleNamespace(
attn_metadata={
"draft.layer": local_draft_metadata,
"target.layer": object(),
}
attn_metadata={"draft.layer": local_draft_metadata, "target.layer": object()},
)
prepared_draft_metadata = global_draft_metadata if replicated_pcp else local_draft_metadata
speculator._prepare_replicated_prefill_attn = MagicMock(
return_value=(
{"draft.layer": prepared_draft_metadata},
None,
)
speculator._build_draft_attn_metadata = MagicMock(
return_value={"draft.layer": global_draft_metadata},
)

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

assert actual == [{"draft.layer": prepared_draft_metadata}]
speculator._prepare_replicated_prefill_attn.assert_called_once_with(
{"draft.layer": local_draft_metadata},
None,
4,
8,
)
expected_metadata = global_draft_metadata if rebuild_metadata else local_draft_metadata
assert actual == [{"draft.layer": expected_metadata}]
assert speculator._build_draft_attn_metadata.call_count == int(rebuild_metadata)
assert local_draft_metadata.decode.actual_seq_lengths_q[-1] == 8
Comment thread
li1how marked this conversation as resolved.


@pytest.mark.skipif(speculator_module.vllm_version_is("0.28.0"), reason="DPSyncState is a main2main interface")
Expand Down
18 changes: 10 additions & 8 deletions vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -496,14 +496,16 @@ def build_draft_attn_metadatas(
}

if is_draft_model_prefill:
prepared_attn_metadata, _ = self._prepare_replicated_prefill_attn(
attn_metadata,
None,
num_reqs_padded,
num_tokens_padded,
)
assert prepared_attn_metadata is not None
return [prepared_attn_metadata]
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]

draft_attn_metadatas = self._init_decode_draft_attn_metadatas(attn_metadata, num_reqs_padded)

Expand Down
Loading