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

120 lines
4.2 KiB
Python

from unittest import mock
import pytest
import torch
from vllm_ascend.device.device_op import A5DeviceAdaptor, BaseDeviceAdaptor
def test_npu_flash_attention_uses_fusion_attention_for_fp32():
query = torch.randn(5, 4, 64, dtype=torch.float32)
key = torch.randn_like(query)
value = torch.randn_like(query)
seq_lens_cpu = torch.tensor([2, 3], dtype=torch.int32)
expected = torch.randn_like(query)
with (
mock.patch(
"vllm_ascend.device.device_op.torch_npu.npu_fusion_attention",
return_value=(expected,),
) as mock_fusion_attention,
mock.patch(
"vllm_ascend.device.device_op.torch_npu._npu_flash_attention_unpad",
create=True,
) as mock_flash_attention,
):
output = BaseDeviceAdaptor.npu_flash_attention(
query=query,
key=key,
value=value,
seq_lens_cpu=seq_lens_cpu,
head_num=4,
scale_value=0.125,
num_kv_heads=4,
)
assert output is expected
mock_flash_attention.assert_not_called()
mock_fusion_attention.assert_called_once()
call_kwargs = mock_fusion_attention.call_args.kwargs
assert call_kwargs["query"] is query
assert call_kwargs["key"] is key
assert call_kwargs["value"] is value
assert call_kwargs["actual_seq_qlen"] == [2, 5]
assert all(isinstance(seq_len, int) for seq_len in call_kwargs["actual_seq_qlen"])
assert call_kwargs["actual_seq_kvlen"] is call_kwargs["actual_seq_qlen"]
assert call_kwargs["head_num"] == 4
assert call_kwargs["scale"] == 0.125
assert call_kwargs["input_layout"] == "TND"
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_npu_flash_attention_uses_unpad_attention_for_low_precision(dtype):
query = torch.randn(5, 4, 64, dtype=dtype)
key = torch.randn_like(query)
value = torch.randn_like(query)
seq_lens_cpu = torch.tensor([2, 3], dtype=torch.int32)
def fake_flash_attention(*, query, key, value, seq_len, scale_value, num_heads, num_kv_heads, out):
out.copy_(query + 1)
with (
mock.patch(
"vllm_ascend.device.device_op.torch_npu.npu_fusion_attention",
) as mock_fusion_attention,
mock.patch(
"vllm_ascend.device.device_op.torch_npu._npu_flash_attention_unpad",
side_effect=fake_flash_attention,
create=True,
) as mock_flash_attention,
):
output = BaseDeviceAdaptor.npu_flash_attention(
query=query,
key=key,
value=value,
seq_lens_cpu=seq_lens_cpu,
head_num=4,
scale_value=0.125,
num_kv_heads=4,
)
mock_fusion_attention.assert_not_called()
mock_flash_attention.assert_called_once()
call_kwargs = mock_flash_attention.call_args.kwargs
assert call_kwargs["query"] is query
assert call_kwargs["key"] is key
assert call_kwargs["value"] is value
assert call_kwargs["seq_len"] is seq_lens_cpu
assert call_kwargs["num_heads"] == 4
assert call_kwargs["num_kv_heads"] == 4
assert call_kwargs["scale_value"] == 0.125
torch.testing.assert_close(output, query + 1)
def test_a5_npu_flash_attention_uses_python_sequence_lengths():
query = torch.randn(5, 4, 64, dtype=torch.float16)
key = torch.randn_like(query)
value = torch.randn_like(query)
seq_lens_cpu = torch.tensor([2, 3], dtype=torch.int32)
expected = torch.randn_like(query)
with mock.patch(
"vllm_ascend.device.device_op.torch_npu.npu_fusion_attention",
return_value=(expected,),
) as mock_fusion_attention:
output = A5DeviceAdaptor.npu_flash_attention(
query=query,
key=key,
value=value,
seq_lens_cpu=seq_lens_cpu,
head_num=4,
scale_value=0.125,
num_kv_heads=4,
)
assert output is expected
call_kwargs = mock_fusion_attention.call_args.kwargs
assert call_kwargs["actual_seq_qlen"] == [2, 5]
assert all(isinstance(seq_len, int) for seq_len in call_kwargs["actual_seq_qlen"])
assert call_kwargs["actual_seq_kvlen"] is call_kwargs["actual_seq_qlen"]