Files
enginex-ascend-910-vllm/tests/ut/ops/test_mm_encoder_attention.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

196 lines
7.4 KiB
Python

from typing import Any
from unittest.mock import MagicMock, patch
import torch
from vllm.config import CompilationConfig, VllmConfig
from vllm.config.vllm import get_cached_compilation_config
from tests.ut.base import TestBase
from vllm_ascend.ops.mm_encoder_attention import (
MAX_PAD_SIZE,
AscendMMEncoderAttention,
)
from vllm_ascend.worker import encoder_acl_graph
from vllm_ascend.worker.encoder_acl_graph import (
get_encoder_graph_params,
set_encoder_forward_context,
set_encoder_graph_params,
)
class FIAMockMixin(TestBase):
captured: dict[str, Any]
def _install_vllm_config_mock(self):
mock_vllm_config = MagicMock(spec=VllmConfig)
mock_vllm_config.compilation_config = CompilationConfig()
patcher = patch(
"vllm.config.vllm.get_current_vllm_config",
return_value=mock_vllm_config,
)
patcher.start()
self.addCleanup(patcher.stop)
get_cached_compilation_config.cache_clear()
self.addCleanup(get_cached_compilation_config.cache_clear)
def _make_layer(self, num_heads=4, num_kv_heads=4, head_size=72, scale=None):
return AscendMMEncoderAttention(
num_heads=num_heads,
head_size=head_size,
scale=scale,
num_kv_heads=num_kv_heads,
)
def _fake_fia(self, **kwargs):
self.captured = {
"mode": "functional",
"q_shape": kwargs["query"].shape,
"input_layout": kwargs["input_layout"],
"actual_seq_lengths": kwargs["actual_seq_lengths"],
}
return torch.zeros_like(kwargs["query"]), None
def _fake_fia_out(self, *, workspace, out, **kwargs):
self.captured = {"mode": "out", "softmax_lse": out[1]}
out[0].zero_()
def _install_fia_mocks(self, *, capture: bool):
self.captured = {}
mock_fia = MagicMock(side_effect=self._fake_fia)
mock_fia.out = self._fake_fia_out
patch_targets: list[tuple[str, Any]] = [
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu.npu_fused_infer_attention_score",
mock_fia,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu._npu_fused_infer_attention_score_get_max_workspace",
MagicMock(return_value=torch.zeros(1)),
),
]
if capture:
self.mock_graph_begin = MagicMock()
self.mock_graph_end = MagicMock(return_value=42)
mock_event = MagicMock()
patch_targets.extend(
[
(
"vllm_ascend.ops.mm_encoder_attention.weak_ref_tensors",
lambda tensors: tensors,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu.npu.current_stream",
MagicMock(return_value=MagicMock()),
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.ExternalEvent",
MagicMock(return_value=mock_event),
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.graph_task_group_begin",
self.mock_graph_begin,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.graph_task_group_end",
self.mock_graph_end,
),
]
)
for target, replacement in patch_targets:
patcher = patch(target, replacement)
patcher.start()
self.addCleanup(patcher.stop)
class TestAscendMMEncoderAttentionEager(FIAMockMixin):
def setUp(self):
self._install_vllm_config_mock()
self._install_fia_mocks(capture=False)
def test_forward_oot_basic(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=128)
bsz, q_len = 2, 4
query = torch.randn(bsz, q_len, layer.num_heads * layer.head_size)
key = query.clone()
value = query.clone()
cu_seqlens = torch.arange(0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32)
out = layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(out.shape, (bsz, q_len, layer.num_heads * layer.head_size))
self.assertEqual(self.captured["mode"], "functional")
self.assertEqual(self.captured["input_layout"], "TND")
def test_forward_oot_seqlens(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
seq_lens = [3, 7, 2]
cu_seqlens = torch.tensor([0, 3, 10, 12], dtype=torch.int32, device="cpu")
max_q_len = max(seq_lens)
query = torch.randn(len(seq_lens), max_q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
out = layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(out.shape, query.shape)
self.assertEqual(self.captured["actual_seq_lengths"], [3, 10, 12])
self.assertEqual(self.captured["q_shape"], (len(seq_lens) * max_q_len, 4, MAX_PAD_SIZE))
class TestAscendMMEncoderAttentionCapture(FIAMockMixin):
def setUp(self):
self._install_vllm_config_mock()
set_encoder_graph_params([2048])
self._install_fia_mocks(capture=True)
def tearDown(self):
encoder_acl_graph._encoder_graph_params = None
encoder_acl_graph._reset_encoder_forward_context()
def test_forward_oot_basic(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
bsz, q_len = 2, 4
query = torch.randn(bsz, q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
cu_seqlens = torch.arange(0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32)
with set_encoder_forward_context(2048, True):
layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
params = get_encoder_graph_params()
self.assertIsNotNone(params)
self.assertEqual(len(params.attn_params[2048]), 1)
self.assertEqual(len(params.handles[2048]), 1)
self.assertEqual(self.captured["mode"], "out")
self.mock_graph_begin.assert_called_once()
self.mock_graph_end.assert_called_once()
def test_forward_oot_seqlens(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
seq_lens = [3, 7, 2]
cu_seqlens = torch.tensor([0, 3, 10, 12], dtype=torch.int32, device="cpu")
max_q_len = max(seq_lens)
query = torch.randn(len(seq_lens), max_q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
captured_lengths: list[Any] = []
def capture_workspace(**kwargs):
captured_lengths.append(kwargs.get("actual_seq_lengths"))
return torch.zeros(1)
with (
patch(
"vllm_ascend.ops.mm_encoder_attention.torch_npu._npu_fused_infer_attention_score_get_max_workspace",
side_effect=capture_workspace,
),
set_encoder_forward_context(2048, True),
):
layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(captured_lengths[-1], [7, 14, 21])