206 lines
6.1 KiB
Python
206 lines
6.1 KiB
Python
from importlib import import_module, util
|
|
|
|
import numpy as np
|
|
import pytest
|
|
import torch
|
|
import torch_npu
|
|
|
|
|
|
def _fa3_available() -> bool:
|
|
try:
|
|
if util.find_spec("flash_attn_npu_v3") is None:
|
|
return False
|
|
mod = import_module("flash_attn_npu_v3")
|
|
return hasattr(mod, "flash_attn_with_kvcache")
|
|
except ImportError:
|
|
return False
|
|
|
|
|
|
def ref_fused_infer_attention(
|
|
query,
|
|
key,
|
|
value,
|
|
block_table,
|
|
block_size,
|
|
actual_seq_lengths_q,
|
|
actual_seq_lengths_kv,
|
|
num_heads,
|
|
num_kv_heads,
|
|
head_size,
|
|
scale,
|
|
attn_mask,
|
|
causal,
|
|
):
|
|
if not causal:
|
|
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
|
query=query,
|
|
key=key,
|
|
value=value,
|
|
block_table=block_table,
|
|
input_layout="TND",
|
|
block_size=block_size,
|
|
actual_seq_lengths=actual_seq_lengths_q,
|
|
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
|
num_key_value_heads=num_kv_heads,
|
|
num_heads=num_heads,
|
|
scale=scale,
|
|
sparse_mode=0,
|
|
)
|
|
else:
|
|
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
|
query=query,
|
|
key=key,
|
|
value=value,
|
|
atten_mask=attn_mask,
|
|
block_table=block_table,
|
|
input_layout="TND",
|
|
block_size=block_size,
|
|
actual_seq_lengths=actual_seq_lengths_q,
|
|
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
|
num_key_value_heads=num_kv_heads,
|
|
num_heads=num_heads,
|
|
scale=scale,
|
|
sparse_mode=3,
|
|
)
|
|
|
|
attn_output = attn_output.view(-1, num_heads, head_size)
|
|
return attn_output
|
|
|
|
|
|
test_cases = [
|
|
# (data_type, batch_size, num_heads, kv_heads, q_seqlen, kv_seqlen, head_size, block_size, is_causal)
|
|
(torch.bfloat16, 1, 1, 1, 1024, 1024, 128, 128, False),
|
|
(torch.bfloat16, 5, 4, 1, 1024, 1024, 128, 128, True),
|
|
(torch.float16, 7, 16, 8, 512, 512, 128, 128, False),
|
|
]
|
|
|
|
|
|
@pytest.mark.skipif(not _fa3_available(), reason="flash_attn_npu_v3 is not installed")
|
|
@pytest.mark.parametrize(
|
|
"data_type, batch_size, num_heads, kv_heads, q_seqlen, kv_seqlen, head_size, block_size, is_causal", test_cases
|
|
)
|
|
def test_fa_custom_ops_tnd(
|
|
data_type, batch_size, num_heads, kv_heads, q_seqlen, kv_seqlen, head_size, block_size, is_causal
|
|
):
|
|
q_min_range = -1.0
|
|
q_max_range = 1.0
|
|
kv_min_range = -1.0
|
|
kv_max_range = 1.0
|
|
block_size = 128
|
|
num_blocks = 64
|
|
|
|
q_sequences = sorted(
|
|
torch.randint(low=1, high=q_seqlen + 1, size=(batch_size,)).tolist(), reverse=False
|
|
) # actual_seq_lengths in fia need in ascending order
|
|
kv_sequences = [torch.randint(low=q, high=kv_seqlen + 1, size=(1,)).item() for q in q_sequences]
|
|
|
|
t_q_sum = sum(q_sequences)
|
|
|
|
query = (q_min_range + (q_max_range - q_min_range) * torch.rand(t_q_sum, num_heads, head_size)).to(data_type).npu()
|
|
|
|
key_cache = None
|
|
value_cache = None
|
|
block_tables = []
|
|
|
|
key_cache = (
|
|
(kv_min_range + (kv_max_range - kv_min_range) * torch.rand(num_blocks, block_size, kv_heads, head_size))
|
|
.to(data_type)
|
|
.npu()
|
|
)
|
|
value_cache = (
|
|
(kv_min_range + (kv_max_range - kv_min_range) * torch.rand(num_blocks, block_size, kv_heads, head_size))
|
|
.to(data_type)
|
|
.npu()
|
|
)
|
|
max_num_blocks_per_seq = (kv_seqlen + block_size - 1) // block_size
|
|
for i in range(batch_size):
|
|
block_table = [max_num_blocks_per_seq * i + j for j in range(max_num_blocks_per_seq)]
|
|
block_tables.append(block_table)
|
|
block_tables = torch.tensor(block_tables, dtype=torch.int32).npu()
|
|
|
|
q_seqlen_list = q_sequences
|
|
kv_seqlen_list = kv_sequences
|
|
|
|
scale = 1.0 / (head_size**0.5)
|
|
window_size_left = -1
|
|
window_size_right = -1
|
|
is_rotary_interleaved = False
|
|
num_splits = 0
|
|
kv_seqlen_list = torch.tensor(kv_seqlen_list, dtype=torch.int32).npu()
|
|
rotary_cos = None
|
|
rotary_sin = None
|
|
cache_batch_idx = None
|
|
leftpad_k = None
|
|
new_q_seqlen_list = None
|
|
|
|
new_q_seqlen_list = [0]
|
|
pre_seq_sum = 0
|
|
for i in range(batch_size):
|
|
pre_seq_sum += q_seqlen_list[i]
|
|
new_q_seqlen_list.append(pre_seq_sum)
|
|
new_q_seqlen_list = torch.tensor(new_q_seqlen_list, dtype=torch.int32).npu()
|
|
|
|
from flash_attn_npu_v3 import flash_attn_with_kvcache # type: ignore[import-not-found]
|
|
|
|
out_out = flash_attn_with_kvcache(
|
|
query,
|
|
key_cache,
|
|
value_cache,
|
|
None,
|
|
None,
|
|
None,
|
|
rotary_cos=rotary_cos,
|
|
rotary_sin=rotary_sin,
|
|
cache_seqlens=kv_seqlen_list,
|
|
cache_batch_idx=cache_batch_idx,
|
|
cache_leftpad=leftpad_k,
|
|
page_table=block_tables,
|
|
cu_seqlens_q=new_q_seqlen_list,
|
|
cu_seqlens_k_new=None,
|
|
max_seqlen_q=q_seqlen,
|
|
rotary_seqlens=None,
|
|
q_descale=None,
|
|
k_descale=None,
|
|
v_descale=None,
|
|
softmax_scale=None,
|
|
causal=is_causal,
|
|
window_size=[window_size_left, window_size_right],
|
|
attention_chunk=0,
|
|
softcap=0.0,
|
|
rotary_interleaved=is_rotary_interleaved,
|
|
scheduler_metadata=None,
|
|
num_splits=num_splits,
|
|
pack_gqa=None,
|
|
sm_margin=0,
|
|
return_softmax_lse=False,
|
|
)
|
|
|
|
ref_out = torch.empty((t_q_sum, num_heads, head_size), dtype=data_type)
|
|
|
|
if is_causal:
|
|
attn_mask = torch.triu(torch.ones(2048, 2048), diagonal=1).to(torch.int8).to(query.device)
|
|
else:
|
|
attn_mask = None
|
|
|
|
q_cumsum = torch.tensor(np.cumsum(q_sequences), dtype=torch.int32, device=query.device)
|
|
|
|
ref_out = ref_fused_infer_attention(
|
|
query,
|
|
key_cache.permute(0, 2, 1, 3).contiguous(), # [num_blocks, num_kv_heads, block_size, head_size]
|
|
value_cache.permute(0, 2, 1, 3).contiguous(),
|
|
block_tables,
|
|
block_size,
|
|
q_cumsum,
|
|
kv_seqlen_list,
|
|
num_heads,
|
|
kv_heads,
|
|
head_size,
|
|
scale,
|
|
attn_mask,
|
|
is_causal,
|
|
)
|
|
|
|
rtol = 1e-2
|
|
atol = 1e-2
|
|
torch.testing.assert_close(out_out.cpu(), ref_out.cpu(), rtol=rtol, atol=atol)
|