Skip to content
Open
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
48 changes: 48 additions & 0 deletions tests/ut/attention/a2/test_attention_cp.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand Down Expand Up @@ -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)
71 changes: 71 additions & 0 deletions tests/ut/attention/a2/test_mla_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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()
Expand Down
62 changes: 61 additions & 1 deletion tests/ut/attention/test_sfa_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -44,14 +45,19 @@ 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,
[],
vllm_config,
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,)

Expand Down Expand Up @@ -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
Loading
Loading