From c14993202b111e5c496dc4ffa8ce946e226f4fe7 Mon Sep 17 00:00:00 2001 From: shenzhao Date: Fri, 11 Sep 2026 10:29:23 +0800 Subject: [PATCH 1/2] fix(attention): isolate PCP global RoPE runtime buffers Keep global KV-update RoPE results separate from rank-local attention metadata while preserving fixed-capacity buffers and stable graph replay addresses. Add numerical, multi-group, capacity and address-stability regression coverage. Signed-off-by: shenzhao --- tests/ut/attention/test_dsa_v1.py | 61 ++++++++++++++ tests/ut/ops/test_rope_dsv4.py | 84 +++++++++++++++++++ .../attention/context_parallel/dsa_cp.py | 1 + vllm_ascend/attention/dsa_v1.py | 3 + vllm_ascend/ops/rope_dsv4.py | 15 ++++ 5 files changed, 164 insertions(+) create mode 100644 tests/ut/ops/test_rope_dsv4.py diff --git a/tests/ut/attention/test_dsa_v1.py b/tests/ut/attention/test_dsa_v1.py index 1b7fdd9d689e..bb751db49332 100644 --- a/tests/ut/attention/test_dsa_v1.py +++ b/tests/ut/attention/test_dsa_v1.py @@ -53,6 +53,7 @@ AscendIndexerMetadata, IndexerOverlapPlan, ) +from vllm_ascend.ops import rope_dsv4 from vllm_ascend.worker.device_metadata import ( DeviceMetadataStage, DeviceMetadataTask, @@ -1667,6 +1668,66 @@ def test_dsa_backend_selects_pcp_and_rejects_legacy_cp(): get_backend_cls() +def test_pcp_builders_keep_global_rope_separate_from_reused_local_metadata(monkeypatch): + state = rope_dsv4.RopeGlobalState() + monkeypatch.setattr(rope_dsv4, "_ROPE_STATE", state) + angles = torch.arange(32 * 4, dtype=torch.float32).reshape(32, 1, 1, 4) + state.full_rope_cache["config"] = (angles.cos(), angles.sin()) + state.registry_summary["config"] = {"default"} + state.layer_info["layer"] = ("config", ["default"]) + state.runtime_buffer["config"] = {"default": (torch.empty(16, 1, 1, 4), torch.empty(16, 1, 1, 4))} + config = _make_vllm_config() + config.parallel_config.prefill_context_parallel_size = 2 + with patch("vllm_ascend.attention.context_parallel.dsa_cp.get_pcp_group") as pcp_group: + pcp_group.return_value.rank_in_group = 0 + builders = [ + AscendDSAPCPMetadataBuilder(_make_kv_cache_spec(4), ["layer"], config, torch.device("cpu")) + for _ in range(2) + ] + for builder in builders: + builder.decode_threshold = builder._global_metadata_builder.decode_threshold = 16 + monkeypatch.setattr(builder, "build_req_metadata", MagicMock()) + monkeypatch.setattr(builder._global_metadata_builder, "build_req_metadata", MagicMock()) + + def build(builder, positions, shared): + offsets = torch.tensor([0, len(positions)], dtype=torch.int32) + common = SimpleNamespace( + num_reqs=1, + num_actual_tokens=len(positions), + num_input_tokens=len(positions), + max_query_len=len(positions), + context_parallel_metadata=None, + query_start_loc=offsets, + query_start_loc_cpu=offsets, + positions=positions, + seq_lens=torch.tensor([32]), + _seq_lens_cpu=torch.tensor([32]), + block_table_tensor=torch.tensor([[0]], dtype=torch.int32), + attn_state=MagicMock(), + ) + AscendDSAMetadataBuilder.build(builder, 0, common, common_ratio_to_sas_metadata=shared) + + addresses = {} + for count in (4, 16, 6): + global_pos, local_pos = torch.arange(count), torch.arange(count - 1) + 5 + local_cache = {} + global_caches = [] + for index, builder in enumerate(builders): + global_cache = {} + build(builder._global_metadata_builder, global_pos, global_cache) + global_caches.append(global_cache) + build(builder, local_pos, local_cache) + for key, full in zip(("cos", "sin"), state.full_rope_cache["config"]): + local = local_cache[key]["layer"] + torch.testing.assert_close(local, full[local_pos], rtol=0, atol=0) + assert local.data_ptr() == addresses.setdefault(("local", key), local.data_ptr()) + for cached in global_caches: + torch.testing.assert_close(cached[key]["layer"], full[global_pos], rtol=0, atol=0) + global_tensor = global_cache[key]["layer"] + assert global_tensor.data_ptr() != local.data_ptr() + assert global_tensor.data_ptr() == addresses.setdefault((index, key), global_tensor.data_ptr()) + + def test_pcp_metadata_builds_from_manager_global_view(): """Build rank-local metadata from the manager's scheduler-global view.""" builder = AscendDSAPCPMetadataBuilder.__new__(AscendDSAPCPMetadataBuilder) diff --git a/tests/ut/ops/test_rope_dsv4.py b/tests/ut/ops/test_rope_dsv4.py new file mode 100644 index 000000000000..d9e511e74cf4 --- /dev/null +++ b/tests/ut/ops/test_rope_dsv4.py @@ -0,0 +1,84 @@ +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch + +from vllm_ascend.ops import rope_dsv4 + + +@pytest.fixture +def rope_state(monkeypatch): + state = rope_dsv4.RopeGlobalState() + monkeypatch.setattr(rope_dsv4, "_ROPE_STATE", state) + for config, width in (("small", 4), ("large", 8)): + angles = torch.arange(32 * width, dtype=torch.float32).reshape(32, 1, 1, width) + state.full_rope_cache[config] = (angles.cos(), angles.sin()) + state.registry_summary[config] = {"default", "compressed"} + state.runtime_buffer[config] = {} + state.spec_runtime_buffer[config] = {} + for group in state.registry_summary[config]: + state.layer_info[f"{config}.{group}"] = (config, [group]) + state.runtime_buffer[config][group] = (torch.empty(16, 1, 1, width), torch.empty(16, 1, 1, width)) + state.spec_runtime_buffer[config][group] = ( + [torch.empty(16, 1, 1, width)], + [torch.empty(16, 1, 1, width)], + ) + return state + + +def _assert_rope(state, result, positions): + for config, cache in state.full_rope_cache.items(): + for group, pos in positions.items(): + for proxy, full in zip(result, cache): + torch.testing.assert_close(proxy[f"{config}.{group}"], full[pos], rtol=0, atol=0) + + +def test_global_local_rope_buffers_survive_interleaved_cache_groups(rope_state): + # Each PCP cache-group builder owns its global buffers, while local + # metadata reuses the first group's result for the rest of the step. + global_buffers = [{}, {}, {}] + addresses = {} + for count in (2, 16, 5, 16, 2): + global_pos = {"default": torch.arange(count), "compressed": torch.arange(count) + 1} + local_pos = {group: pos.flip(0)[: max(1, count - 1)] + 3 for group, pos in global_pos.items()} + global_results = [] + for group_index, buffers in enumerate(global_buffers): + global_result = rope_dsv4.get_cos_and_sin_dsa(global_pos, use_cache=True, runtime_buffer=buffers) + global_results.append(global_result) + if group_index == 0: + local_result = rope_dsv4.get_cos_and_sin_dsa(local_pos, use_cache=True) + # A later global build must preserve the already-created local view. + _assert_rope(rope_state, local_result, local_pos) + for result in global_results: + _assert_rope(rope_state, result, global_pos) + for config, groups in buffers.items(): + for group, pair in groups.items(): + for index, buf in enumerate(pair): + default = rope_state.runtime_buffer[config][group][index] + assert buf.shape == default.shape + assert buf.data_ptr() != default.data_ptr() + key = (group_index, config, group, index) + assert buf.data_ptr() == addresses.setdefault(key, buf.data_ptr()) + assert len(set(addresses.values())) == len(addresses) + + +@pytest.mark.parametrize("draft_index", [None, 1]) +def test_uncached_rope_does_not_allocate_caller_buffers(rope_state, draft_index): + buffers = {} + positions = {"default": torch.tensor([3, 1, 7])} + result = rope_dsv4.get_cos_and_sin_dsa(positions, draft_index=draft_index, runtime_buffer=buffers) + _assert_rope(rope_state, result, positions) + assert buffers == {} + + +def test_draft_rope_keeps_existing_speculative_buffers(rope_state): + buffers = {} + positions = {"default": torch.tensor([3, 1, 7])} + result = rope_dsv4.get_cos_and_sin_dsa(positions, use_cache=True, draft_index=1, runtime_buffer=buffers) + _assert_rope(rope_state, result, positions) + for config in rope_state.full_rope_cache: + for index, proxy in enumerate(result): + assert proxy[f"{config}.default"].data_ptr() == ( + rope_state.spec_runtime_buffer[config]["default"][index][0].data_ptr() + ) + assert buffers == {} diff --git a/vllm_ascend/attention/context_parallel/dsa_cp.py b/vllm_ascend/attention/context_parallel/dsa_cp.py index 89675fef4457..d4b760b81b72 100644 --- a/vllm_ascend/attention/context_parallel/dsa_cp.py +++ b/vllm_ascend/attention/context_parallel/dsa_cp.py @@ -2161,6 +2161,7 @@ def __init__( vllm_config, device, metadata_cls=dsa_v1.AscendDSAMetadata, + rope_runtime_buffer={}, ) self._pcp_world_size = vllm_config.parallel_config.prefill_context_parallel_size self._pcp_rank = get_pcp_group().rank_in_group diff --git a/vllm_ascend/attention/dsa_v1.py b/vllm_ascend/attention/dsa_v1.py index 6c0432e21618..5ab72372c073 100644 --- a/vllm_ascend/attention/dsa_v1.py +++ b/vllm_ascend/attention/dsa_v1.py @@ -593,8 +593,10 @@ def __init__( device: torch.device, metadata_cls: type[AscendDSAMetadata] | None = None, supports_dcp_with_varlen: bool = False, + rope_runtime_buffer: dict[str, dict[str, tuple[torch.Tensor, torch.Tensor]]] | None = None, ): self.kv_cache_spec = kv_cache_spec + self.rope_runtime_buffer = rope_runtime_buffer self.metadata_cls = metadata_cls if metadata_cls is not None else AscendDSAMetadata self.vllm_config = vllm_config self.model_config = vllm_config.model_config @@ -842,6 +844,7 @@ def build( cos, sin = get_cos_and_sin_dsa( input_positions, use_cache=self.num_prefills == 0, + runtime_buffer=self.rope_runtime_buffer, ) self.common_ratio_to_sas_metadata["cos"] = cos self.common_ratio_to_sas_metadata["sin"] = sin diff --git a/vllm_ascend/ops/rope_dsv4.py b/vllm_ascend/ops/rope_dsv4.py index f193d3a2e362..4973df40607c 100644 --- a/vllm_ascend/ops/rope_dsv4.py +++ b/vllm_ascend/ops/rope_dsv4.py @@ -84,6 +84,7 @@ def get_cos_and_sin_dsa( positions: torch.Tensor | dict[str, torch.Tensor], use_cache: bool = False, draft_index: int | None = None, + runtime_buffer: dict[str, dict[str, tuple[torch.Tensor, torch.Tensor]]] | None = None, ): if isinstance(positions, torch.Tensor): pos_map = {"default": positions} @@ -113,6 +114,20 @@ def get_cos_and_sin_dsa( if group_buffers is None: continue + # PCP's global KV update and rank-local attention use different + # positions in the same step. Keep caller-owned buffers stable + # for graph replay without aliasing the default local cache. + if runtime_buffer is not None and draft_index is None: + isolated_groups = runtime_buffer.setdefault(config_key, {}) + if group_name not in isolated_groups: + # Allocate the full registered capacity even when the + # first call is a small warmup/capture batch. + isolated_groups[group_name] = ( + torch.empty_like(group_buffers[0]), + torch.empty_like(group_buffers[1]), + ) + group_buffers = isolated_groups[group_name] + buf_cos, buf_sin = group_buffers num_tokens = pos_tensor.size(0) From ca169f36aef921203a9857cb15cbf12a4ef695a1 Mon Sep 17 00:00:00 2001 From: shenzhao Date: Fri, 11 Sep 2026 10:39:14 +0800 Subject: [PATCH 2/2] test: annotate PCP RoPE regression buffers Add concrete dictionary types required by the repository mypy checks. Production code and the original e2e test remain byte-identical to the fix commit. Signed-off-by: shenzhao --- tests/ut/attention/test_dsa_v1.py | 6 +++--- tests/ut/ops/test_rope_dsv4.py | 8 ++++---- 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/tests/ut/attention/test_dsa_v1.py b/tests/ut/attention/test_dsa_v1.py index bb751db49332..a99cc6321826 100644 --- a/tests/ut/attention/test_dsa_v1.py +++ b/tests/ut/attention/test_dsa_v1.py @@ -1707,13 +1707,13 @@ def build(builder, positions, shared): ) AscendDSAMetadataBuilder.build(builder, 0, common, common_ratio_to_sas_metadata=shared) - addresses = {} + addresses: dict[tuple[int | str, str], int] = {} for count in (4, 16, 6): global_pos, local_pos = torch.arange(count), torch.arange(count - 1) + 5 - local_cache = {} + local_cache: dict[str, Any] = {} global_caches = [] for index, builder in enumerate(builders): - global_cache = {} + global_cache: dict[str, Any] = {} build(builder._global_metadata_builder, global_pos, global_cache) global_caches.append(global_cache) build(builder, local_pos, local_cache) diff --git a/tests/ut/ops/test_rope_dsv4.py b/tests/ut/ops/test_rope_dsv4.py index d9e511e74cf4..56b7e7e79971 100644 --- a/tests/ut/ops/test_rope_dsv4.py +++ b/tests/ut/ops/test_rope_dsv4.py @@ -36,8 +36,8 @@ def _assert_rope(state, result, positions): def test_global_local_rope_buffers_survive_interleaved_cache_groups(rope_state): # Each PCP cache-group builder owns its global buffers, while local # metadata reuses the first group's result for the rest of the step. - global_buffers = [{}, {}, {}] - addresses = {} + global_buffers: list[dict[str, dict[str, tuple[torch.Tensor, torch.Tensor]]]] = [{}, {}, {}] + addresses: dict[tuple[int, str, str, int], int] = {} for count in (2, 16, 5, 16, 2): global_pos = {"default": torch.arange(count), "compressed": torch.arange(count) + 1} local_pos = {group: pos.flip(0)[: max(1, count - 1)] + 3 for group, pos in global_pos.items()} @@ -64,7 +64,7 @@ def test_global_local_rope_buffers_survive_interleaved_cache_groups(rope_state): @pytest.mark.parametrize("draft_index", [None, 1]) def test_uncached_rope_does_not_allocate_caller_buffers(rope_state, draft_index): - buffers = {} + buffers: dict[str, dict[str, tuple[torch.Tensor, torch.Tensor]]] = {} positions = {"default": torch.tensor([3, 1, 7])} result = rope_dsv4.get_cos_and_sin_dsa(positions, draft_index=draft_index, runtime_buffer=buffers) _assert_rope(rope_state, result, positions) @@ -72,7 +72,7 @@ def test_uncached_rope_does_not_allocate_caller_buffers(rope_state, draft_index) def test_draft_rope_keeps_existing_speculative_buffers(rope_state): - buffers = {} + buffers: dict[str, dict[str, tuple[torch.Tensor, torch.Tensor]]] = {} positions = {"default": torch.tensor([3, 1, 7])} result = rope_dsv4.get_cos_and_sin_dsa(positions, use_cache=True, draft_index=1, runtime_buffer=buffers) _assert_rope(rope_state, result, positions)