Skip to content
Closed
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
61 changes: 61 additions & 0 deletions tests/ut/attention/test_dsa_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
AscendIndexerMetadata,
IndexerOverlapPlan,
)
from vllm_ascend.ops import rope_dsv4
from vllm_ascend.worker.device_metadata import (
DeviceMetadataStage,
DeviceMetadataTask,
Expand Down Expand Up @@ -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: 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: dict[str, Any] = {}
global_caches = []
for index, builder in enumerate(builders):
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)
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)
Expand Down
84 changes: 84 additions & 0 deletions tests/ut/ops/test_rope_dsv4.py
Original file line number Diff line number Diff line change
@@ -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: 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()}
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: 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)
assert buffers == {}


def test_draft_rope_keeps_existing_speculative_buffers(rope_state):
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)
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 == {}
1 change: 1 addition & 0 deletions vllm_ascend/attention/context_parallel/dsa_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 3 additions & 0 deletions vllm_ascend/attention/dsa_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
15 changes: 15 additions & 0 deletions vllm_ascend/ops/rope_dsv4.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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)

Expand Down
Loading