Skip to content

Commit a0d0671

Browse files
authored
[BugFix]Add check for empty _ATTN_KEYS_BUFFER (vllm-project#11214)
### What this PR does / why we need it? Adds a check for empty ``_ATTN_KEYS_BUFFER`` and corresponding tests. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@a30addc --------- Signed-off-by: zhiyu-wa <1959864813@qq.com>
1 parent d22f3e6 commit a0d0671

3 files changed

Lines changed: 62 additions & 2 deletions

File tree

tests/e2e/pull_request/four_card/test_graph_mode.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -366,6 +366,11 @@
366366
"capture_mem_tolerance": 1.5,
367367
}
368368

369+
CASE_DS_ACLGRAPH_ENPU = {
370+
**CASE_DS_ACLGRAPH,
371+
"env_vars": {"ENPU_ENABLE": "true"},
372+
}
373+
369374
# inherit from tests/e2e/pull_request/utils.py::compare_logprobs
370375
ATOL = 0.0689
371376

@@ -493,6 +498,9 @@ def _run_worker_process(
493498
}
494499
)
495500

501+
for key, value in cur_case.get("env_vars", {}).items():
502+
os.environ[key] = str(value)
503+
496504
# Apply hooks and run inference
497505
with _install_spies(metrics):
498506
short_prompts = cur_case["prompts"]["short"]
@@ -595,7 +603,7 @@ def check_capture_mem(capture_mem, baseline_capture_mem=0.2, capture_mem_toleran
595603

596604

597605
@wait_until_npu_memory_free(0.7)
598-
@pytest.mark.parametrize("cur_case", [CASE_QWEN_ACLGRAPH, CASE_DS_ACLGRAPH])
606+
@pytest.mark.parametrize("cur_case", [CASE_QWEN_ACLGRAPH, CASE_DS_ACLGRAPH, CASE_DS_ACLGRAPH_ENPU])
599607
def test_aclgraph(cur_case: dict, monkeypatch: pytest.MonkeyPatch):
600608
# Counter doesn't work in default "spawn" mode
601609
metrics = None

tests/ut/attention/a2/test_attention_v1.py

Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33

44
import torch
55

6+
import vllm_ascend.attention.attention_v1 as attn_module
67
from tests.ut.base import TestBase
78
from vllm_ascend.attention.attention_v1 import (
89
AscendAttentionBackend,
@@ -605,3 +606,54 @@ def test_get_kvcomp_params_decode_hamming(self, mock_hamming, mock_reshape):
605606
# Verify the result is recorded in hamming_output_records
606607
self.assertTrue(torch.equal(kvcomp_metadata.hamming_output, torch.ones(2, 5)))
607608
self.assertIsNotNone(kvcomp_metadata.seq_lens_for_hamming)
609+
610+
@patch("torch.npu.stream")
611+
@patch("torch.npu.graph_task_update_begin")
612+
@patch("torch.npu.graph_task_update_end")
613+
@patch("torch_npu.npu_fused_infer_attention_score")
614+
@patch("vllm_ascend.attention.attention_v1.get_graph_params")
615+
@patch("vllm_ascend.attention.attention_v1._EXTRA_CTX")
616+
@patch("vllm_ascend.attention.attention_v1.using_paged_attention", return_value=False)
617+
@patch("vllm_ascend.attention.attention_v1.needs_layer_aware_fia_graph_replay", return_value=False)
618+
@patch("vllm_ascend.attention.attention_v1._ATTN_KEYS_BUFFER", new=[])
619+
def test_update_graph_params(
620+
self,
621+
mock_needs_layer_aware_fia_graph_replay,
622+
mock_using_paged_attention,
623+
mock_EXTRA_CTX,
624+
mock_get_graph_params,
625+
mock_fia,
626+
mock_graph_task_update_end,
627+
mock_graph_task_update_begin,
628+
mock_stream,
629+
):
630+
"""Test behavior when _ATTN_KEYS_BUFFER is [] after dummy_run."""
631+
632+
mock_EXTRA_CTX.sinks = False
633+
mock_EXTRA_CTX.is_draft_model = False
634+
635+
param: list[MagicMock | None] = [MagicMock()] * 21
636+
param[16] = None
637+
param[20] = None
638+
639+
mock_get_graph_params.return_value.attn_params = {1: [tuple(param)] * 3}
640+
mock_get_graph_params.return_value.handles = {1: [MagicMock()] * 3}
641+
mock_get_graph_params.return_value.events = {1: [MagicMock()] * 3}
642+
643+
attn_metadata_keys = [
644+
"model.layers.10.self_attn.attn",
645+
"model.layers.2.self_attn.attn",
646+
"model.layers.5.self_attn.attn",
647+
]
648+
forward_context = MagicMock()
649+
forward_context.attn_metadata = {key: MagicMock() for key in attn_metadata_keys}
650+
# breakpoint()
651+
self.impl.update_graph_params(self.mock_stream, forward_context, 1, self.mock_vllm_config)
652+
653+
expected = [
654+
"model.layers.2.self_attn.attn",
655+
"model.layers.5.self_attn.attn",
656+
"model.layers.10.self_attn.attn",
657+
]
658+
self.assertEqual(attn_module._ATTN_KEYS_BUFFER, expected)
659+
self.assertEqual(mock_fia.out.call_count, 3)

vllm_ascend/attention/attention_v1.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -644,7 +644,7 @@ def update_graph_params(
644644
global _ATTN_KEYS_BUFFER
645645
if attn_keys_length == 0:
646646
return
647-
if _ATTN_KEYS_BUFFER is None or len(_ATTN_KEYS_BUFFER) != attn_keys_length:
647+
if not _ATTN_KEYS_BUFFER or len(_ATTN_KEYS_BUFFER) != attn_keys_length:
648648
import regex as re
649649

650650
def extract_layer_index(key: str) -> int:

0 commit comments

Comments
 (0)