From c75b33d34e7fca0a041a5c72230ba445075fcef9 Mon Sep 17 00:00:00 2001 From: weiguihua2 Date: Mon, 14 Sep 2026 16:41:44 +0800 Subject: [PATCH] [BugFix] resolve shape mismatch in DCP PD-disaggregated recomputation (#16487) Fix shape mismatches when DCP is enabled with PD-disaggregated recomputation. Short recomputation requests, including last-token recomputation, are now correctly classified as decode requests in the attention, MLA, and SFA metadata builders. Cache DCP state during initialization and use the builder's configuration so metadata construction works outside the current-config context. No API or configuration changes. This fixes failures in DCP-enabled PD-disaggregated recomputation. - Added regression tests covering PD recomputation, mixed query lengths, and behavior when the override does not apply. - Added coverage for metadata construction without an active current-config context. - vLLM main: https://github.com/vllm-project/vllm/commit/a97dacb7106ee49f39f3d1fc6ae1800ff724e01d --------- Signed-off-by: weiguihua2 (cherry picked from commit cab37196df6366edd0eb38a0e35f15c9dc1e4f15) --- tests/ut/attention/a2/test_attention_cp.py | 48 +++++++++++++ tests/ut/attention/a2/test_mla_v1.py | 71 +++++++++++++++++++ tests/ut/attention/test_sfa_cp.py | 62 +++++++++++++++- .../context_parallel/attention_cp.py | 15 +++- .../attention/context_parallel/sfa_cp.py | 8 ++- vllm_ascend/attention/mla_v1.py | 9 ++- 6 files changed, 207 insertions(+), 6 deletions(-) diff --git a/tests/ut/attention/a2/test_attention_cp.py b/tests/ut/attention/a2/test_attention_cp.py index 1f0bca345b4b..4f3f1d06f21c 100644 --- a/tests/ut/attention/a2/test_attention_cp.py +++ b/tests/ut/attention/a2/test_attention_cp.py @@ -1,6 +1,10 @@ # SPDX-License-Identifier: Apache-2.0 +from types import SimpleNamespace +from unittest.mock import patch + import numpy as np +import pytest import torch from vllm_ascend.attention.attention_v1 import ( @@ -73,3 +77,47 @@ def test_dcp_partial_attention_merge_matches_weighted_reference() -> None: torch.testing.assert_close(output, torch.tensor([[[4.0, 6.0]]])) torch.testing.assert_close(merged_lse, torch.tensor([[[np.log(4.0)]]], dtype=torch.float32)) + + +@pytest.mark.parametrize( + "is_consumer,is_producer,recompute", [(True, False, True), (True, False, False), (False, True, True)] +) +@pytest.mark.parametrize("query_lens", [[1, 1], [3, 3], [3, 5]]) +def test_dcp_split_uses_builder_config_without_current_context(is_consumer, is_producer, recompute, query_lens): + config = SimpleNamespace( + kv_transfer_config=SimpleNamespace(is_kv_consumer=is_consumer, is_kv_producer=is_producer), + ) + with ( + patch( + "vllm_ascend.attention.context_parallel.attention_cp.DCPMetadataBuilderMixin.__init__", return_value=None + ), + patch("vllm_ascend.attention.context_parallel.attention_cp.enable_dcp", return_value=True) as dcp, + ): + builder = AscendAttentionDCPMetadataBuilder() + dcp.assert_called_once_with() + builder.vllm_config = config + builder.decode_threshold = 3 + query_start_loc = torch.tensor([0, query_lens[0], sum(query_lens)], dtype=torch.int32) + common = SimpleNamespace( + context_parallel_metadata=None, + max_query_len=max(query_lens), + num_reqs=2, + num_actual_tokens=sum(query_lens), + query_start_loc_cpu=query_start_loc, + is_prefilling=torch.ones(2, dtype=torch.bool), + ) + with ( + patch("vllm.config.get_current_vllm_config", side_effect=AssertionError("no current context")), + patch( + "vllm_ascend.utils.get_ascend_config", + return_value=SimpleNamespace(scheduler_config=SimpleNamespace(recompute_scheduler_enable=recompute)), + ), + patch( + "vllm_ascend.attention.context_parallel.attention_cp.enable_dcp", + side_effect=AssertionError("use cached DCP state"), + ), + ): + actual = builder._split_decodes_and_prefills(common) + num_decodes = sum(q <= 3 for q in query_lens) if is_consumer and not is_producer and recompute else 0 + num_decode_tokens = sum(query_lens[:num_decodes]) + assert actual == (num_decodes, 2 - num_decodes, num_decode_tokens, sum(query_lens) - num_decode_tokens) diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index e490d262982b..5d271948472c 100755 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -298,6 +298,10 @@ def test_ascend_mla_metadata_default(self): class TestAscendMLAMetadataBuilder(TestBase): def setUp(self): + dcp_patcher = patch("vllm_ascend.attention.mla_v1.enable_dcp", return_value=False) + dcp_patcher.start() + self.addCleanup(dcp_patcher.stop) + # Mock parent class __init__ to avoid complex initialization, # but still set the essential attributes that child class needs def mock_parent_init( @@ -556,6 +560,10 @@ def test_pad_actual_seq_lens_q_mtp_enable_pad_with_padding(self): class TestAscendMLAMetadataBuilderBuild(TestBase): def setUp(self): + dcp_patcher = patch("vllm_ascend.attention.mla_v1.enable_dcp", return_value=False) + dcp_patcher.start() + self.addCleanup(dcp_patcher.stop) + # Mock parent class __init__ to avoid complex initialization, # but still set the essential attributes that child class needs def mock_parent_init( @@ -758,6 +766,69 @@ def test_build_decode_only_metadata(self, mock_get_cos_and_sin_mla): self.assertEqual(metadata.head_dim, self.kv_cache_spec.head_size) self.assertEqual(metadata.decode.seq_lens_device.data_ptr(), seq_lens_device.data_ptr()) + # PD recomputes the last prompt token (N-1 computed). Metadata building + # runs outside set_current_vllm_config, unlike DCP manager initialization. + self.mock_vllm_config.parallel_config.decode_context_parallel_size = 16 + self.mock_vllm_config.kv_transfer_config = SimpleNamespace(is_kv_consumer=True, is_kv_producer=False) + common_attn_metadata.is_prefilling = torch.ones(3, dtype=torch.bool) + # DCP runs populate context_parallel_metadata, so the release-branch + # `is None` gate is False and classification hinges on the override. + common_attn_metadata.context_parallel_metadata = SimpleNamespace( + query_lens_cpu=torch.tensor([1, 1, 1], dtype=torch.int32), + max_query_len=1, + ) + with patch("vllm_ascend.attention.mla_v1.enable_dcp", return_value=True) as mock_enable_dcp: + builder = AscendMLAMetadataBuilder( + self.kv_cache_spec, ["layer_0", "layer_1"], self.mock_vllm_config, self.mock_device + ) + mock_enable_dcp.assert_called_once_with() + self.assertTrue(builder.dcp_enabled) + ascend_config = SimpleNamespace(scheduler_config=SimpleNamespace(recompute_scheduler_enable=True)) + with ( + patch("vllm.config.get_current_vllm_config", side_effect=AssertionError("no current context")), + patch("vllm_ascend.utils.get_ascend_config", return_value=ascend_config), + patch( + "vllm_ascend.attention.mla_v1.enable_dcp", + side_effect=AssertionError("DCP state must be cached during initialization"), + ), + ): + metadata = builder.build(0, common_attn_metadata) + self.assertEqual(metadata.num_decodes, 3) + self.assertEqual(metadata.num_prefills, 0) + self.assertEqual(metadata.num_decode_tokens, 3) + self.assertIsNone(metadata.prefill) + self.assertEqual(metadata.decode.seq_lens_list, [4, 5, 6]) + + # With DCP enabled but the recompute scheduler off, the override must + # not fire: short extends stay prefills. + ascend_config.scheduler_config.recompute_scheduler_enable = False + with ( + patch("vllm.config.get_current_vllm_config", side_effect=AssertionError("no current context")), + patch("vllm_ascend.utils.get_ascend_config", return_value=ascend_config), + patch.object(builder, "build_prefill_metadata", return_value=MagicMock()), + ): + metadata = builder.build(0, common_attn_metadata) + self.assertEqual(metadata.num_decodes, 0) + self.assertEqual(metadata.num_prefills, 3) + self.assertEqual(metadata.num_decode_tokens, 0) + common_attn_metadata.context_parallel_metadata = None + + # Without DCP, preserve the original classification even on a PD consumer: + # the DCP-only override must never be evaluated. + self.mock_vllm_config.parallel_config.decode_context_parallel_size = 1 + builder.dcp_enabled = False + with ( + patch("vllm.config.get_current_vllm_config", side_effect=AssertionError("no current context")), + patch( + "vllm_ascend.attention.mla_v1.is_pd_decode_recompute_scheduler_enabled", + side_effect=AssertionError("DCP-only override must not run without DCP"), + ), + ): + metadata = builder.build(0, common_attn_metadata) + self.assertEqual(metadata.num_decodes, 3) + self.assertEqual(metadata.num_prefills, 0) + self.assertEqual(metadata.num_decode_tokens, 3) + @patch("vllm_ascend.attention.mla_v1.get_cos_and_sin_mla") def test_build_decode_metadata_without_disable_padded_drafter_batch(self, mock_get_cos_and_sin_mla): common_attn_metadata = MagicMock() diff --git a/tests/ut/attention/test_sfa_cp.py b/tests/ut/attention/test_sfa_cp.py index 0971aaba6049..a7b29423b07a 100644 --- a/tests/ut/attention/test_sfa_cp.py +++ b/tests/ut/attention/test_sfa_cp.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from unittest.mock import patch +import pytest import torch from vllm_ascend.attention.context_parallel.common_cp import DCPMetadataBuilderMixin @@ -44,7 +45,10 @@ def fake_base_init(self, *args, **kwargs) -> None: model_config=SimpleNamespace(max_model_len=1024), ) - with patch.object(DCPMetadataBuilderMixin, "__init__", new=fake_base_init): + with ( + patch("vllm_ascend.attention.context_parallel.sfa_cp.enable_dcp", return_value=True) as dcp, + patch.object(DCPMetadataBuilderMixin, "__init__", new=fake_base_init), + ): builder = AscendSFADCPMetadataBuilder( kv_cache_spec, [], @@ -52,6 +56,8 @@ def fake_base_init(self, *args, **kwargs) -> None: torch.device("cpu"), ) + dcp.assert_called_once_with() + assert builder.dcp_enabled assert builder.block_table_replicated_view_buf.shape == (5, 8) assert builder.arange_buffer.shape == (8,) @@ -119,3 +125,57 @@ def test_sfa_dcp_updates_dsa_cp_local_slot_mapping_with_padding() -> None: dsa_cp_context.slot_mapping_cp, torch.tensor([12, 13, -1], dtype=torch.int32), ) + + +@pytest.mark.parametrize( + "is_consumer,is_producer,recompute", [(True, False, True), (True, False, False), (False, True, True)] +) +@pytest.mark.parametrize("query_lens", [[1, 1], [3, 3], [3, 5]]) +def test_sfa_dcp_split_uses_builder_config_without_current_context(is_consumer, is_producer, recompute, query_lens): + builder = _make_builder() + builder.dcp_enabled = True + builder.decode_threshold = 3 + builder.vllm_config = SimpleNamespace( + kv_transfer_config=SimpleNamespace(is_kv_consumer=is_consumer, is_kv_producer=is_producer), + ) + builder.dcp_local_seq_lens_buf = torch.empty(2, dtype=torch.int32) + slots = torch.arange(sum(query_lens), dtype=torch.int64) + blocks = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32) + common = SimpleNamespace( + context_parallel_metadata=None, + max_query_len=max(query_lens), + num_reqs=2, + num_actual_tokens=sum(query_lens), + num_input_tokens=sum(query_lens), + query_start_loc_cpu=torch.tensor([0, query_lens[0], sum(query_lens)], dtype=torch.int32), + is_prefilling=torch.ones(2, dtype=torch.bool), + slot_mapping=slots, + block_table_tensor=blocks, + seq_lens=torch.tensor([10, 20], dtype=torch.int32), + dcp_local_seq_lens=torch.tensor([6, 12], dtype=torch.int32), + ) + metadata = AscendSFADCPMetadata.__new__(AscendSFADCPMetadata) + with ( + patch("vllm.config.get_current_vllm_config", side_effect=AssertionError("no current context")), + patch( + "vllm_ascend.utils.get_ascend_config", + return_value=SimpleNamespace(scheduler_config=SimpleNamespace(recompute_scheduler_enable=recompute)), + ), + patch( + "vllm_ascend.attention.context_parallel.sfa_cp.enable_dcp", + side_effect=AssertionError("use cached DCP state"), + ), + patch.object(builder, "_get_dcp_local_block_table", return_value=blocks), + patch.object(builder, "_build_block_table_replicated_view", return_value=blocks), + patch.object(builder, "_build_slot_mapping_replicated_view", return_value=slots), + patch.object(builder, "_build_compact_kv_gather_metadata", return_value=(torch.arange(4), blocks)) as gather, + patch.object(builder, "_update_dsa_cp_slot_mapping_for_dcp"), + ): + result = builder._build_with_metadata_view(common, lambda: metadata) + num_decodes = sum(q <= 3 for q in query_lens) if is_consumer and not is_producer and recompute else 0 + assert result.num_decodes == num_decodes + assert result.num_prefills == 2 - num_decodes + assert result.num_decode_tokens == sum(query_lens[:num_decodes]) + assert gather.call_count == int(result.num_prefills > 0) + assert common.slot_mapping is slots + assert common.block_table_tensor is blocks diff --git a/vllm_ascend/attention/context_parallel/attention_cp.py b/vllm_ascend/attention/context_parallel/attention_cp.py index 2f78042974ae..2ad8d6122908 100644 --- a/vllm_ascend/attention/context_parallel/attention_cp.py +++ b/vllm_ascend/attention/context_parallel/attention_cp.py @@ -36,6 +36,7 @@ ) from vllm_ascend.attention.utils import ( AscendCommonAttentionMetadata, + enable_dcp, filter_chunked_req_indices, split_decodes_and_prefills, ) @@ -47,7 +48,11 @@ ) from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.memcache_comm_fence import record_attention_compute_start -from vllm_ascend.utils import cp_chunkedprefill_comm_stream, weak_ref_tensors +from vllm_ascend.utils import ( + cp_chunkedprefill_comm_stream, + is_pd_decode_recompute_scheduler_enabled, + weak_ref_tensors, +) @dataclass @@ -94,6 +99,10 @@ class AscendAttentionDCPMetadataBuilder( metadata_cls = AscendAttentionDCPMetadata + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + self.dcp_enabled = enable_dcp() + def _split_decodes_and_prefills( self, common_attn_metadata: AscendCommonAttentionMetadata, @@ -101,7 +110,9 @@ def _split_decodes_and_prefills( return split_decodes_and_prefills( common_attn_metadata, decode_threshold=self.decode_threshold, - treat_short_extends_as_decodes=False, + treat_short_extends_as_decodes=( + self.dcp_enabled and is_pd_decode_recompute_scheduler_enabled(self.vllm_config) + ), ) @staticmethod diff --git a/vllm_ascend/attention/context_parallel/sfa_cp.py b/vllm_ascend/attention/context_parallel/sfa_cp.py index 62f8890b4b56..f18b7c71861d 100644 --- a/vllm_ascend/attention/context_parallel/sfa_cp.py +++ b/vllm_ascend/attention/context_parallel/sfa_cp.py @@ -21,9 +21,10 @@ AscendSFAMetadataBuilder, DSACPContext, ) -from vllm_ascend.attention.utils import AscendCommonAttentionMetadata, split_decodes_and_prefills +from vllm_ascend.attention.utils import AscendCommonAttentionMetadata, enable_dcp, split_decodes_and_prefills from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.distributed.utils import all_gather_async +from vllm_ascend.utils import is_pd_decode_recompute_scheduler_enabled M = TypeVar("M", bound=AscendSFAMetadata) @@ -90,6 +91,7 @@ def __init__( metadata_cls, supports_dcp_with_varlen, ) + self.dcp_enabled = enable_dcp() self.cp_kv_cache_interleave_size = vllm_config.parallel_config.cp_kv_cache_interleave_size assert self.dcp_size > 1, "AscendSFADCPMetadataBuilder requires DCP world size > 1." if self.cp_kv_cache_interleave_size <= 0: @@ -331,7 +333,9 @@ def _build_with_metadata_view( num_decodes, num_prefills, num_decode_tokens, _ = split_decodes_and_prefills( common_attn_metadata, decode_threshold=self.decode_threshold, - treat_short_extends_as_decodes=False, + treat_short_extends_as_decodes=( + self.dcp_enabled and is_pd_decode_recompute_scheduler_enabled(self.vllm_config) + ), ) kv_gather_block_ids = None kv_gather_block_table = None diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index 10e35efb3670..c3653fe0167a 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -54,6 +54,7 @@ ACL_FORMAT_FRACTAL_ND, AscendDeviceType, get_ascend_device_type, + is_pd_decode_recompute_scheduler_enabled, maybe_trans_nz, weak_ref_tensors, ) @@ -409,6 +410,7 @@ def __init__( metadata_cls if metadata_cls is not None else AscendMLAMetadata, supports_dcp_with_varlen, ) + self.dcp_enabled = enable_dcp() scheduler_config = vllm_config.scheduler_config self.block_size = vllm_config.cache_config.block_size @@ -597,7 +599,12 @@ def build( split_decodes_and_prefills( common_attn_metadata, decode_threshold=self.decode_threshold, - treat_short_extends_as_decodes=common_attn_metadata.context_parallel_metadata is None, + treat_short_extends_as_decodes=( + common_attn_metadata.context_parallel_metadata is None + # Only DCP needs the PD last-token recompute override. + # Use the builder's config outside the current-config context. + or (self.dcp_enabled and is_pd_decode_recompute_scheduler_enabled(self.vllm_config)) + ), ) ) self.set_num_actual_tokens(common_attn_metadata)