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/tests/ut/kv_offload/test_mooncake_connector.py b/tests/ut/kv_offload/test_mooncake_connector.py index 6ab3b7d647e0..de7c490f1d70 100644 --- a/tests/ut/kv_offload/test_mooncake_connector.py +++ b/tests/ut/kv_offload/test_mooncake_connector.py @@ -1,3 +1,4 @@ +import math import os import queue import socket @@ -3547,6 +3548,10 @@ def test_start_load_kv_puts_replicated_indexer_on_existing_transfer_port(self): worker.kv_send_thread = None worker.kv_recv_thread = MagicMock() worker._prefill_tp_size = 4 + worker.kv_group2layeridx = { + 0: ({"kv_cache_spec_type": "AscendMLAAttentionSpec"}, [0]), + 1: ({"kv_cache_spec_type": "AscendSFAIndexerCacheSpec"}, [2]), + } worker.remote_port_send_num = {"remote_engine": {31001: {"num": 1, "host": "localhost"}}} worker._get_sfa_replicate_k_block_ids = MagicMock(return_value=(([40],), ([20],))) worker._get_kv_split_metadata = MagicMock( @@ -3558,8 +3563,8 @@ def test_start_load_kv_puts_replicated_indexer_on_existing_transfer_port(self): ) worker._get_group_pulls_metadata = MagicMock( return_value=[ - [[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]], - [[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]], + [[GroupPull(group_id=g, remote_tp_offset=0, num_group_pulls=1) for g in (0, 1)]], + [[GroupPull(group_id=g, remote_tp_offset=0, num_group_pulls=1) for g in (0, 1)]], ] ) worker._get_remote_host_info_by_port = MagicMock(return_value=("localhost", "remote_engine")) @@ -3589,6 +3594,8 @@ def test_start_load_kv_puts_replicated_indexer_on_existing_transfer_port(self): self.assertEqual(add_request_calls[1].kwargs["remote_handshake_port"], 31003) self.assertIsNone(add_request_calls[1].kwargs["local_block_ids_replicate_k"]) self.assertIsNone(add_request_calls[1].kwargs["remote_block_ids_replicate_k"]) + self.assertEqual([pull.group_id for pull in add_request_calls[0].kwargs["group_pulls"]], [0, 1]) + self.assertEqual([pull.group_id for pull in add_request_calls[1].kwargs["group_pulls"]], [0]) def test_get_kv_split_metadata_dp1_remote_port_send_num_uses_absolute_ports(self): self.vllm_config.kv_transfer_config.kv_port = 30000 @@ -3734,6 +3741,223 @@ def test_get_sfa_replicated_indexer_block_ids_when_only_remote_enables_dcp(self) self.assertEqual(local_ids, ([20, 21],)) self.assertEqual(remote_ids, ([20, 21],)) + def test_mrv2_dcp_transfers_follow_global_block_ownership(self): + # Release-branch adaptation: PCP stays a KV shard axis here, so the + # upstream PCP=2 topology does not apply; ownership is exercised for + # pure-DCP geometries (remote_dcp bounded by prefill TP). + for remote_dcp in (2, 8): + for local_dcp in (1, 8): + for rank in range(local_dcp): + with self.subTest(remote_dcp=remote_dcp, local_dcp=local_dcp, rank=rank): + worker = self._build_non_cp_worker() + worker.use_mla = True + worker.use_sfa_sparse = True + worker.enable_sfa_dcp_replicated_indexer = True + worker.tp_size = worker._prefill_tp_size = 8 + worker.tp_rank = worker.dcp_rank = rank + worker.dcp_size = local_dcp + worker.pcp_size = 1 + worker.pcp_rank = 0 + worker.handshake_port = worker.side_channel_port + rank + worker.local_remote_block_port_mapping = {} + worker.remote_port_send_num = {} + worker.block_size_scale = [[1]] + worker.kv_group2layeridx = {0: ({"kv_cache_spec_type": "MLAAttentionSpec"}, [0])} + meta = types.SimpleNamespace( + remote_pcp_size=1, + remote_dcp_size=remote_dcp, + remote_ptp_size=8, + remote_port=30000, + remote_block_ids=(list(range(100, 100 + math.ceil(33 / remote_dcp))),), + local_block_ids=(list(range(200, 200 + math.ceil(33 / local_dcp))),), + local_full_block_ids=(list(range(200, 200 + math.ceil(33 / local_dcp))),), + num_external_tokens=33 * worker.block_size, + num_prompt_blocks=33, + num_computed_tokens=0, + remote_block_size=worker.block_size, + remote_engine_id="pcp_test", + remote_host="localhost", + remote_multi_nodes_meta_mapping={}, + ) + ports, local_ids, remote_ids = worker._get_kv_split_metadata("pcp", cast(ReqMeta, meta)) + seen = [] + for shard_ports, dst, src in zip(ports, local_ids, remote_ids): + self.assertEqual(len(shard_ports), 1) + offset = shard_ports[0] - meta.remote_port + # Release cp_group layout: port = base + dcp_rank + # + dcp_repeat_offset (a multiple of remote_dcp). + source_cp_rank = offset % remote_dcp + for local_block, remote_block in zip(dst[0], src[0]): + global_block = (remote_block - 100) * remote_dcp + source_cp_rank + self.assertEqual(global_block % local_dcp, rank) + self.assertEqual(local_block, 200 + global_block // local_dcp) + seen.append(global_block) + self.assertEqual(sorted(seen), list(range(rank, 33, local_dcp))) + local_index, remote_index = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta)) + self.assertEqual(local_index, (list(range(200 * local_dcp, 200 * local_dcp + 33)),)) + self.assertEqual(remote_index, (list(range(100 * remote_dcp, 100 * remote_dcp + 33)),)) + pulls = worker._get_dcp_shard_pulls(ports, 8, 30000, remote_pcp_size=1) + self.assertTrue( + all(p.prefill_pp_rank == 0 for shard in pulls for group in shard for p in group) + ) + + def test_sfa_replicated_indexer_cp_ratio(self): + for remote_cp, local_cp in ((8, 2), (2, 8), (4, 4), (6, 4), (0, 4), (4, 0)): + with self.subTest(remote_cp=remote_cp, local_cp=local_cp): + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) + worker.enable_sfa_dcp_replicated_indexer = True + worker.use_sfa_sparse = True + # Release-branch _get_sfa_replicate_k_block_ids reads + # vllm_config.model_config directly. + worker.vllm_config = types.SimpleNamespace(model_config=None) + worker.pcp_size = 1 + worker.dcp_size = local_cp + worker.block_size = 16 + meta = types.SimpleNamespace( + remote_pcp_size=1, + remote_dcp_size=remote_cp, + remote_block_ids=([10, 11, 12, 13],), + local_block_ids=([20, 21, 22, 23],), + local_full_block_ids=([20, 21, 22, 23],), + num_external_tokens=128, + num_prompt_blocks=8, + num_computed_tokens=0, + ) + if (remote_cp, local_cp) in ((6, 4), (0, 4), (4, 0)): + with self.assertRaisesRegex(AssertionError, "one divisible by the other"): + worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta)) + else: + local_ids, remote_ids = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta)) + self.assertEqual(local_ids, (list(range(20 * local_cp, 20 * local_cp + 8)),)) + self.assertEqual(remote_ids, (list(range(10 * remote_cp, 10 * remote_cp + 8)),)) + + def test_decode_only_dcp_checks_each_cache_group_type(self): + for spec_type in ("MLAAttentionSpec", "AscendMLAAttentionSpec", "FullAttentionSpec"): + with self.subTest(spec_type=spec_type): + worker = self._build_non_cp_worker() + worker.dcp_size = 2 + worker.dcp_rank = 0 + worker.block_size_scale = [[1], [1]] + worker.kv_group2layeridx = { + 0: ({"kv_cache_spec_type": "MLAAttentionSpec", "kv_cache_group_id": 0}, [0]), + 1: ({"kv_cache_spec_type": spec_type, "kv_cache_group_id": 1}, [1]), + } + meta = types.SimpleNamespace( + remote_pcp_size=1, + remote_dcp_size=1, + remote_ptp_size=1, + remote_port=30000, + remote_block_size=16, + local_block_ids=([20], [30]), + local_full_block_ids=([20], [30]), + remote_block_ids=([100, 101], [200, 201]), + num_prompt_blocks=2, + num_computed_tokens=0, + ) + if spec_type in ("MLAAttentionSpec", "AscendMLAAttentionSpec"): + _, local_ids, remote_ids = worker._get_kv_split_metadata("r", cast(ReqMeta, meta)) + self.assertEqual(local_ids, [([20], [30])]) + self.assertEqual(remote_ids, [([100], [200])]) + else: + with self.assertRaisesRegex(NotImplementedError, f"{spec_type} in transfer group 1"): + worker._get_kv_split_metadata("r", cast(ReqMeta, meta)) + + def test_decode_only_dcp_keeps_kda_state_and_tp_owners(self): + for state_first in (False, True): + for prefill_tp in (2, 4): + for rank in range(2): + for prompt_blocks in (1, 5): + with self.subTest( + state_first=state_first, prefill_tp=prefill_tp, rank=rank, prompt_blocks=prompt_blocks + ): + worker = self._build_non_cp_worker() + worker._is_hma_required = True + worker.use_mla = True + worker.dcp_size = worker.tp_size = 2 + worker.dcp_rank = worker.tp_rank = rank + worker._prefill_tp_size = prefill_tp + worker._prefill_pp_size = 1 + worker.block_size_scale = [[2], [1]] + mla = (0, ({"kv_cache_spec_type": "MLAAttentionSpec", "kv_cache_group_id": 0}, [0])) + state = (1, ({"kv_cache_spec_type": "MambaSpec", "kv_cache_group_id": 1}, [1])) + worker.kv_group2layeridx = dict([state, mla] if state_first else [mla, state]) + meta = types.SimpleNamespace( + remote_pcp_size=1, + remote_dcp_size=1, + remote_ptp_size=prefill_tp, + remote_port=30000, + remote_block_size=16, + num_computed_tokens=0, + num_prompt_blocks=prompt_blocks, + remote_block_ids=(list(range(100, 100 + prompt_blocks)), [200]), + local_block_ids=([20, 21, 22], [30]), + # State uses its transfer IDs, not the attention prefix-cache view. + local_full_block_ids=([20, 21, 22], [99]), + ) + ports, local_ids, remote_ids = worker._get_kv_split_metadata("k3", cast(ReqMeta, meta)) + selected = range(rank, prompt_blocks, 2) + self.assertEqual( + local_ids[0][0], [2 * (20 + g // 2) + k for g in selected for k in range(2)] + ) + self.assertEqual(remote_ids[0][0], [2 * (100 + g) + k for g in selected for k in range(2)]) + self.assertEqual(local_ids[0][1], [30]) + self.assertEqual(remote_ids[0][1], [200]) + _, owners = worker._get_hybrid_remote_rank_group_pulls("k3", prefill_tp) + pulls = worker._get_group_pulls_metadata("k3", ports, prefill_tp, 30000, 1, 1) + self.assertEqual(pulls, [[owners[port - 30000] for port in ports[0]]]) + state_ports = { + port + for port, groups in zip(ports[0], pulls[0]) + if any(group.group_id == 1 for group in groups) + } + self.assertEqual( + state_ports, + {30000 + rank * (prefill_tp // 2) + offset for offset in range(prefill_tp // 2)}, + ) + + def test_sfa_decode_only_dcp_maps_global_blocks_to_each_rank(self): + # Release-branch adaptation: the decode-only DCP path is gated to + # PCP == 1 on both sides (PCP replica selection needs the MRV2 + # replica semantics that this branch does not carry). + for rank, remote_pcp_size in ((rank, pcp) for rank in range(8) for pcp in (1,)): + for prompt_blocks, prefix_blocks in ((1, 0), (17, 0), (17, 9)): + with self.subTest(rank=rank, pcp=remote_pcp_size, prompt=prompt_blocks, prefix=prefix_blocks): + worker = self._build_non_cp_worker() + worker.use_sfa_sparse = True + worker.enable_sfa_dcp_replicated_indexer = True + worker.dcp_size = 8 + worker.dcp_rank = rank + worker._get_selected_pcp_rank = MagicMock(return_value=remote_pcp_size - 1) + worker.block_size_scale = [[2], [8]] + worker.kv_group2layeridx = { + 0: ({"kv_cache_spec_type": "MLAAttentionSpec", "kv_cache_group_id": 0}, [0]), + 1: ({"kv_cache_spec_type": "AscendSFAIndexerCacheSpec", "kv_cache_group_id": 0}, [1]), + } + meta = types.SimpleNamespace( + remote_pcp_size=remote_pcp_size, + remote_dcp_size=1, + remote_ptp_size=1, + remote_port=30000, + remote_block_ids=(list(range(100, 100 + prompt_blocks)),), + local_block_ids=([20, 21, 22],), + local_full_block_ids=([20, 21, 22],), + num_external_tokens=(prompt_blocks - prefix_blocks) * 16, + num_prompt_blocks=prompt_blocks, + num_computed_tokens=prefix_blocks * 16, + remote_block_size=16, + remote_engine_id="sfa_p1_d8", + remote_host="localhost", + remote_multi_nodes_meta_mapping={}, + ) + ports, local_ids, remote_ids = worker._get_kv_split_metadata("r", cast(ReqMeta, meta)) + selected = [g for g in range(prefix_blocks, prompt_blocks) if g % 8 == rank] + self.assertEqual(ports, [[30000 + remote_pcp_size - 1]]) + self.assertEqual(local_ids, [([2 * (20 + g // 8) + k for g in selected for k in range(2)], [])]) + self.assertEqual(remote_ids, [([2 * (100 + g) + k for g in selected for k in range(2)], [])]) + local_index, remote_index = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta)) + self.assertEqual(local_index, ([160 + g for g in range(prefix_blocks, prompt_blocks)],)) + self.assertEqual(remote_index, ([100 + g for g in range(prefix_blocks, prompt_blocks)],)) + def test_get_sfa_replicated_indexer_block_ids_requires_full_blocks_for_prefix(self): worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) worker.enable_sfa_dcp_replicated_indexer = True 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) diff --git a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py index 3b831c74b081..1d3a6ad18481 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py +++ b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py @@ -2668,6 +2668,8 @@ def _local_kernel_ids_for_shard( # block_idx-th pulled block maps to global prompt block (in P-units): # global_p_block = (shard_first_p_block + block_idx) * Rcp + shard_cp_rank global_p_block = (shard_first_p_block + block_idx) * remote_cp_size + shard_cp_rank + if (global_p_block // block_size_ratio) % local_cp_size != rank_first_d_block % local_cp_size: + continue if remote_block_size > self.block_size: # Bp > Bd (only supported when D-side has no CP): one P-block spans multiple # D-blocks, so walk it kernel by kernel via the absolute token offset within @@ -2714,6 +2716,12 @@ def _group_compress_ratio(group_spec): def _get_kv_cache_group_id(group_idx: int, group_spec: dict[str, Any]) -> int: return group_spec.get("kv_cache_group_id", group_idx) + def _get_kernel_block_scale(self, layer_indices: list[int]) -> int: + """Kernel block scale for logical-to-tensor block expansion.""" + if layer_indices and layer_indices[0] < len(self.block_size_scale) and self.block_size_scale[layer_indices[0]]: + return self.block_size_scale[layer_indices[0]][0] + return 1 + def _get_kernel_block_ids(self, layer_indices, meta, group_idx, group_spec): """No-CP per-group block ids at kernel granularity: (local, remote). @@ -2770,6 +2778,10 @@ def _get_local_remote_cp_params(self, meta: ReqMeta): Also validates that P/D block sizes are compatible under D-side CP. """ remote_block_size = meta.remote_block_size or self.block_size + # MRV2's DCP group already spans PCP; PCP is not another KV shard axis. + # Keep PCP in the CP layout on this release branch: without the MRV2 + # "DCP spans PCP" invariant (main #15809), PCP is still a KV shard axis + # here. For pcp == 1 these values equal main's DCP-only forms. local_cp_rank = self.dcp_rank + self.pcp_rank * self.dcp_size local_cp_size = self.dcp_size * self.pcp_size remote_cp_size = meta.remote_pcp_size * meta.remote_dcp_size @@ -2791,6 +2803,63 @@ def _get_local_remote_cp_params(self, meta: ReqMeta): r_blk = self.block_size // remote_block_size if self.block_size > remote_block_size else 1 return remote_block_size, local_cp_rank, local_cp_size, remote_cp_size, r_blk + def _get_decode_only_dcp_metadata( + self, + req_id: str, + meta: ReqMeta, + prefill_tp_size: int, + ) -> tuple[list[list[int]], list[BlockIds], list[BlockIds]]: + assert (meta.remote_block_size or self.block_size) == self.block_size, ( + "Decode-only DCP requires equal P/D block sizes." + ) + if self._is_hma_required: + chosen_rank_list, _ = self._get_hybrid_remote_rank_group_pulls(req_id, prefill_tp_size) + else: + chosen_rank_list = self._get_remote_rank(req_id, prefill_tp_size) + pcp_offset = self._get_selected_pcp_rank(req_id, meta.remote_pcp_size) * prefill_tp_size + remote_handshake_port_list = [[x + meta.remote_port + pcp_offset for x in chosen_rank_list]] + use_transfer_group_block_ids = transfer_groups_need_independent_block_ids( + self.kv_group2layeridx, self.block_size_scale + ) + local_block_ids: list[list[int]] = [ + [] for _ in (self.kv_group2layeridx if use_transfer_group_block_ids else meta.local_block_ids) + ] + remote_block_ids: list[list[int]] = [ + [] for _ in (self.kv_group2layeridx if use_transfer_group_block_ids else meta.remote_block_ids) + ] + for group_idx, (group_spec, layer_indices) in self.kv_group2layeridx.items(): + spec_type = group_spec["kv_cache_spec_type"] + group_id = self._get_kv_cache_group_id(group_idx, group_spec) + block_id_idx = group_idx if use_transfer_group_block_ids else group_id + if spec_type == "MambaSpec": + # KDA keeps the full sequence state for this TP rank's heads. + local_block_ids[block_id_idx], remote_block_ids[block_id_idx] = self._get_kernel_block_ids( + layer_indices, meta, group_idx, group_spec + ) + continue + if spec_type == "AscendSFAIndexerCacheSpec": + # The full indexer cache is transferred separately. + continue + if spec_type not in ("MLAAttentionSpec", "AscendMLAAttentionSpec"): + raise NotImplementedError( + f"Decode-only DCP does not support cache type {spec_type} " + f"in transfer group {group_idx} (layer indices {layer_indices})." + ) + local_blocks = (meta.local_full_block_ids or meta.local_block_ids)[group_id] + remote_blocks = meta.remote_block_ids[group_id] + first_block = meta.num_computed_tokens // self.block_size + first_block += (self.dcp_rank - first_block) % self.dcp_size + # P owns the full sequence; D rank r owns r, r + DCP, ... . + global_blocks = range(first_block, min(meta.num_prompt_blocks, len(remote_blocks)), self.dcp_size) + scale = self._get_kernel_block_scale(layer_indices) + local_block_ids[block_id_idx] = self._expand_block_ids( + [local_blocks[block // self.dcp_size] for block in global_blocks], scale + ) + remote_block_ids[block_id_idx] = self._expand_block_ids( + [remote_blocks[block] for block in global_blocks], scale + ) + return remote_handshake_port_list, [tuple(local_block_ids)], [tuple(remote_block_ids)] + def _get_kv_split_metadata( self, req_id: str, @@ -2823,6 +2892,15 @@ def _get_kv_split_metadata( """ prefill_tp_size: int = meta.remote_ptp_size if meta.remote_ptp_size is not None else self._prefill_tp_size + # The decode-only-DCP path assumes main's replica semantics for PCP; + # this release branch keeps PCP as a shard axis, so take it only when + # neither side uses PCP. + is_decode_only_dcp = ( + self.dcp_size > 1 and meta.remote_dcp_size == 1 and self.pcp_size == 1 and meta.remote_pcp_size == 1 + ) + if is_decode_only_dcp: + return self._get_decode_only_dcp_metadata(req_id, meta, prefill_tp_size) + if meta.remote_pcp_size * meta.remote_dcp_size * self.pcp_size * self.dcp_size == 1: if self._is_hma_required: chosen_rank_list, _ = self._get_hybrid_remote_rank_group_pulls(req_id, prefill_tp_size) @@ -2850,7 +2928,17 @@ def _get_kv_split_metadata( ) def context_parallel_parameters_check(): - assert (meta.remote_pcp_size * meta.remote_dcp_size) % (self.pcp_size * self.dcp_size) == 0 + if self.pcp_size > 1 or meta.remote_pcp_size > 1: + # Legacy PCP transfers keep the parent's one-direction rule: + # source selection and completion accounting for a larger + # local CP size are only implemented for pure-DCP setups. + assert remote_cp_size % local_cp_size == 0, ( + f"P CP size({remote_cp_size}) must be divisible by D CP size({local_cp_size}) when PCP is used." + ) + else: + assert remote_cp_size % local_cp_size == 0 or local_cp_size % remote_cp_size == 0, ( + f"P/D CP sizes must be divisible in either direction, got P={remote_cp_size}, D={local_cp_size}." + ) if not (self.use_mla or self.use_sparse): p_node_heads_per_rank = math.ceil(self.num_key_value_heads / prefill_tp_size) d_node_heads_per_rank = math.ceil(self.num_key_value_heads / self.tp_size) @@ -2930,7 +3018,7 @@ def get_local_remote_block_port_mappings(): for p_idx, p_port in enumerate(p_cp_group): # When Bd == Bp, r_blk = 1, which degenerates to the original `p_idx % Lcp` rule. # When Bd = r * Bp, all blocks of P CP rank q are mapped to D rank `(q // r) % Lcp`. - if (p_idx // r_blk) % len(d_cp_group) == d_idx: + if (p_idx // r_blk - d_idx) % min(len(d_cp_group), len(p_cp_group)) == 0: p_port_remote_list.append(p_port) local_remote_block_port_mappings[d_port].append(p_port_remote_list) @@ -3055,10 +3143,8 @@ def _set_hma_shared_port(prefill_tp_size, meta, remote_handshake_port_list, req_ ), 0, ) - assert math.ceil(num_external_blocks / (self.pcp_size * self.dcp_size)) == len( - meta.local_block_ids[sequence_group_idx] - ), ( - f"num_external_blocks({num_external_blocks}), cp_size({self.pcp_size * self.dcp_size}), " + assert math.ceil(num_external_blocks / local_cp_size) == len(meta.local_block_ids[sequence_group_idx]), ( + f"num_external_blocks({num_external_blocks}), cp_size({local_cp_size}), " f"local_block_ids_len ({len(meta.local_block_ids[sequence_group_idx])})" ) assert meta.num_prompt_blocks >= num_external_blocks_p, ( @@ -3085,7 +3171,7 @@ def _set_hma_shared_port(prefill_tp_size, meta, remote_handshake_port_list, req_ for cp_rank, block_num in enumerate(remote_block_nums_all): # When r_blk = 1, it degrades to the original cp_rank % Lcp rule. - if (cp_rank // r_blk) % local_cp_size == local_cp_rank: + if (cp_rank // r_blk - local_cp_rank) % min(local_cp_size, remote_cp_size) == 0: if last_block_location == cp_rank: final_block_idx = len(remote_block_nums) remote_block_nums.append(block_num) @@ -3170,6 +3256,13 @@ def _set_hma_shared_port(prefill_tp_size, meta, remote_handshake_port_list, req_ remote_logical = list( meta.remote_block_ids[kv_cache_group_id][remote_first : remote_first + num_blocks_to_pull] ) + if local_cp_size > remote_cp_size: + remote_logical = [ + block_id + for offset, block_id in enumerate(remote_logical) + if ((remote_first + offset) * remote_cp_size + shard_cp_rank) // r_blk % local_cp_size + == local_cp_rank + ] kernel_remote = self._expand_block_ids(remote_logical, remote_scale) kernel_local = self._local_kernel_ids_for_shard( remote_first, @@ -3205,9 +3298,8 @@ def _set_hma_shared_port(prefill_tp_size, meta, remote_handshake_port_list, req_ return remote_handshake_port_list, local_block_ids_list, remote_block_ids_list - def _get_cp_shard_pulls(self, remote_handshake_port_list, prefill_tp_size, remote_base_port, remote_pcp_size): - # CP case: `group_pulls` is derived from `port` (which already includes the random selection result), - # eliminating the need for a table lookup. + def _get_dcp_shard_pulls(self, remote_handshake_port_list, prefill_tp_size, remote_base_port, remote_pcp_size): + """Build group pulls from the selected ports of each DCP shard.""" mamba_num = prefill_tp_size // self.tp_size attn_num = self._get_tp_num_need_pulls(prefill_tp_size) attn_gids = [ @@ -3224,9 +3316,8 @@ def _get_cp_shard_pulls(self, remote_handshake_port_list, prefill_tp_size, remot for port_idx, port in enumerate(ports): pulls = [] port_tp = (port - remote_base_port) % prefill_tp_size - # PCP and PP are mutually exclusive; when PCP > 1, pp_rank is always 0. - pp_rank = 0 if remote_pcp_size > 1 else (port - remote_base_port) // prefill_tp_size - # The first attn_num ports of each shard (i.e., the original ports with randomly substituted TPs). + pp_rank = (port - remote_base_port) // (prefill_tp_size * remote_pcp_size) + # Attention uses the leading ports selected for each DCP shard. if port_idx < attn_num: pulls += [ GroupPull( @@ -3288,17 +3379,23 @@ def _get_group_pulls_metadata( this pull is the final pull for the group. The final-pull flag is used by the receiver to decide when group reformatting can run. """ + # PCP transfers still take the shard-pulls path on this release branch + # (PCP remains a shard axis here; see _validate_and_get_cp_layout). cp_transfer = remote_pcp_size * remote_dcp_size * self.pcp_size * self.dcp_size > 1 + # Decode-only DCP selects its ports via the hybrid rank table + # (_get_decode_only_dcp_metadata), so its pulls must come from the + # same table; the shard-pulls builder would drop attention pulls for + # later pipeline stages. + is_decode_only_dcp = self.dcp_size > 1 and remote_pcp_size * remote_dcp_size * self.pcp_size == 1 if self._is_hma_required: - if not cp_transfer: + if not cp_transfer or is_decode_only_dcp: # Non-CP case: port = base + chosen_rank, which has a one-to-one correspondence # with the table keys, maintaining the original logic. _, rank_group_pulls = self._get_hybrid_remote_rank_group_pulls(req_id, prefill_tp_size) return [[rank_group_pulls[p - remote_base_port] for p in ports] for ports in remote_handshake_port_list] - # CP case: `group_pulls` is derived from `port` (which already includes the random selection result), - # eliminating the need for a table lookup. - return self._get_cp_shard_pulls( + # The DCP path has already selected the source ports for each shard. + return self._get_dcp_shard_pulls( remote_handshake_port_list, prefill_tp_size, remote_base_port, remote_pcp_size ) @@ -3483,14 +3580,26 @@ def _get_sfa_replicate_k_block_ids( f"Got remote groups={len(meta.remote_block_ids)}, local groups={len(meta.local_block_ids)}." ) + # PCP stays a shard axis on this release branch (see _validate_and_get_cp_layout). remote_cp_size = meta.remote_pcp_size * meta.remote_dcp_size local_cp_size = self.pcp_size * self.dcp_size - if local_cp_size == 0 or remote_cp_size % local_cp_size != 0: + if (self.pcp_size > 1 or meta.remote_pcp_size > 1) and ( + local_cp_size <= 0 or remote_cp_size <= 0 or remote_cp_size % local_cp_size != 0 + ): + # Legacy PCP transfers keep the parent's one-direction rule. raise AssertionError( f"SFA replicate-K expects remote cp size({remote_cp_size}) to be divisible by " f"local cp size({local_cp_size})." ) - + if ( + local_cp_size <= 0 + or remote_cp_size <= 0 + or (remote_cp_size % local_cp_size != 0 and local_cp_size % remote_cp_size != 0) + ): + raise AssertionError( + f"SFA replicate-K requires positive P/D CP sizes with one divisible by the other, " + f"got P CP={remote_cp_size}, D CP={local_cp_size}." + ) num_prefix_cached_blocks = min(meta.num_computed_tokens // self.block_size, meta.num_prompt_blocks) num_external_blocks = meta.num_prompt_blocks - num_prefix_cached_blocks num_external_blocks_from_tokens = math.ceil(meta.num_external_tokens / self.block_size) @@ -3586,7 +3695,7 @@ def start_load_kv(self, metadata: MooncakeConnectorMetadata): meta.remote_dcp_size, ) - for pcp_dcp_rank, remote_ports in enumerate(remote_handshake_port_list): + for shard_idx, remote_ports in enumerate(remote_handshake_port_list): for remote_tp_offset, remote_handshake_port in enumerate(remote_ports): assert self.kv_recv_thread is not None remote_host, remote_engine_id = self._get_remote_host_info_by_port( @@ -3611,22 +3720,32 @@ def start_load_kv(self, metadata: MooncakeConnectorMetadata): if replicate_k_transfer_port is not None and remote_handshake_port == replicate_k_transfer_port else None ) + group_pulls = group_pulls_list[shard_idx][remote_tp_offset] + if has_replicate_k_blocks and remote_handshake_port != replicate_k_transfer_port: + # The indexer is replicated, not an attention DCP shard. + # Other ports must not overwrite its full-cache transfer. + group_pulls = [ + pull + for pull in group_pulls + if self.kv_group2layeridx[pull.group_id][0]["kv_cache_spec_type"] + != "AscendSFAIndexerCacheSpec" + ] self.kv_recv_thread.add_request( request_id=req_id, remote_request_id=remote_req_id, - local_block_ids=local_block_ids_list[pcp_dcp_rank], - remote_block_ids=remote_block_ids_list[pcp_dcp_rank], - group_pulls=group_pulls_list[pcp_dcp_rank][remote_tp_offset], + local_block_ids=local_block_ids_list[shard_idx], + remote_block_ids=remote_block_ids_list[shard_idx], + group_pulls=group_pulls, remote_engine_id=remote_engine_id, remote_host=remote_host, remote_handshake_port=remote_handshake_port, remote_port_send_num=remote_port_send_num, num_computed_tokens=meta.num_computed_tokens, all_task_done=( - pcp_dcp_rank == len(remote_handshake_port_list) - 1 + shard_idx == len(remote_handshake_port_list) - 1 and remote_tp_offset == len(remote_ports) - 1 ), - shard_idx=pcp_dcp_rank, + shard_idx=shard_idx, remote_block_size=meta.remote_block_size, local_block_ids_replicate_k=local_block_ids_replicate_k_for_port, remote_block_ids_replicate_k=remote_block_ids_replicate_k_for_port, @@ -3658,6 +3777,14 @@ def _get_tp_num_need_pulls(self, prefill_tp_size: int | None) -> int: tp_num_need_pulls = num_d_block_heads // num_p_block_heads return tp_num_need_pulls + @staticmethod + def _get_selected_pcp_rank(req_id: str, pcp_size: int) -> int: + if pcp_size == 1: + return 0 + # P and D use the P request ID to select the same replica, independently + # of TP routing. + return random.Random(string_to_int64_hash(f"pcp:{req_id}")).randrange(pcp_size) + def _get_remote_host_info_by_port( self, base_port: int,