Files
project_6/verify_flash_attn.py
project6-dev 9f02200ede diag: verify_flash_attn.py — test flash_attn_func correctness + perf for head_dim=256
flash_attn_func WORKS with head_dim=256 on BI-V100!
This is the path to 10-50x attention speedup.
Tests: correctness vs ref, GQA, long seq, varlen, paged decode, perf.
2026-08-13 04:29:49 +00:00

184 lines
7.7 KiB
Python

#!/usr/bin/env python3
"""Verify ixformer.flash_attn_func for Qwen3.5 prefill attention.
flash_attn_func works with head_dim=256 on BI-V100!
Now test correctness vs PyTorch ref and benchmark on real prefill lengths.
Also test flash_attn_varlen_func (used by vllm for variable-length batching)
and vllm_single_query_cached_kv_attention (used for decode with KV cache).
"""
import sys
import time
import torch
import torch.nn.functional as F
def pytorch_attention_ref(q, k, v, causal=True):
"""(batch, seqlen, nheads, headdim) format."""
q = q.transpose(1, 2) # (B, H, S, D)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
scale = q.shape[-1] ** -0.5
attn = torch.matmul(q.float() * scale, k.float().transpose(-2, -1))
if causal and q.shape[-2] > 1:
L, S = q.shape[-2], k.shape[-2]
mask = torch.triu(torch.ones(L, S, device=q.device, dtype=torch.bool),
diagonal=S - L + 1)
attn = attn.masked_fill(mask, float('-inf'))
attn = torch.softmax(attn, dim=-1)
out = torch.matmul(attn, v.float())
return out.transpose(1, 2).to(q.dtype) # back to (B, S, H, D)
def main():
print("=" * 60)
print("BI-V100 flash_attn_func verification for Qwen3.5")
print("=" * 60)
import ixformer
# Qwen3.5 full attention dims (TP=4):
# num_heads=4, num_kv_heads=1, head_dim=256
num_heads = 4
num_kv_heads = 1
head_dim = 256
# --- Test 1: Correctness with GQA (different q/kv heads) ---
print("\n--- Test 1: Correctness (GQA: q_heads=4, kv_heads=1) ---")
for seq_len in [1, 4, 16, 64, 128, 256]:
torch.manual_seed(42)
q = torch.randn(1, seq_len, num_heads, head_dim,
device="cuda", dtype=torch.float16)
k = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
v = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
try:
out = ixformer.flash_attn_func(q, k, v, causal=(seq_len > 1))
# For ref, expand kv heads to match q
k_exp = k.expand(-1, -1, num_heads, -1)
v_exp = v.expand(-1, -1, num_heads, -1)
ref = pytorch_attention_ref(q, k_exp, v_exp, causal=(seq_len > 1))
diff = (out.float() - ref.float()).abs().max().item()
has_nan = out.isnan().any().item()
status = "PASS" if diff < 0.05 and not has_nan else "FAIL"
print(f" seq={seq_len:4d}: diff={diff:.6f} nan={has_nan} {status}")
except Exception as e:
print(f" seq={seq_len:4d}: EXCEPTION: {e}")
# --- Test 2: Longer sequences (actual prefill lengths) ---
print("\n--- Test 2: Long sequence prefill ---")
for seq_len in [512, 1024, 2048, 4096]:
q = torch.randn(1, seq_len, num_heads, head_dim,
device="cuda", dtype=torch.float16)
k = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
v = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
try:
out = ixformer.flash_attn_func(q, k, v, causal=True)
has_nan = out.isnan().any().item()
print(f" seq={seq_len:5d}: shape={out.shape} nan={has_nan}")
except Exception as e:
print(f" seq={seq_len:5d}: EXCEPTION: {e}")
# --- Test 3: flash_attn_varlen_func (variable length, used by vllm) ---
print("\n--- Test 3: flash_attn_varlen_func ---")
if hasattr(ixformer, 'flash_attn_varlen_func'):
for seq_len in [64, 256, 1024]:
q = torch.randn(seq_len, num_heads, head_dim,
device="cuda", dtype=torch.float16)
k = torch.randn(seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
v = torch.randn(seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
cu_seqlens = torch.tensor([0, seq_len], device="cuda", dtype=torch.int32)
try:
out = ixformer.flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens,
seq_len, seq_len, causal=True)
has_nan = out.isnan().any().item()
print(f" varlen seq={seq_len:5d}: shape={out.shape} nan={has_nan}")
except Exception as e:
print(f" varlen seq={seq_len:5d}: EXCEPTION: {e}")
# --- Test 4: Performance ---
print("\n--- Test 4: Performance flash_attn_func vs PyTorch ---")
for seq_len in [64, 256, 1024]:
q = torch.randn(1, seq_len, num_heads, head_dim,
device="cuda", dtype=torch.float16)
k = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
v = torch.randn(1, seq_len, num_kv_heads, head_dim,
device="cuda", dtype=torch.float16)
k_exp = k.expand(-1, -1, num_heads, -1).contiguous()
v_exp = v.expand(-1, -1, num_heads, -1).contiguous()
# Warmup
for _ in range(5):
ixformer.flash_attn_func(q, k, v, causal=True)
pytorch_attention_ref(q, k_exp, v_exp, causal=True)
torch.cuda.synchronize()
N = 20
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(N):
ixformer.flash_attn_func(q, k, v, causal=True)
torch.cuda.synchronize()
ix_ms = (time.perf_counter() - t0) / N * 1000
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(N):
pytorch_attention_ref(q, k_exp, v_exp, causal=True)
torch.cuda.synchronize()
pt_ms = (time.perf_counter() - t0) / N * 1000
print(f" seq={seq_len:5d}: flash={ix_ms:.2f}ms pytorch={pt_ms:.2f}ms "
f"speedup={pt_ms/ix_ms:.1f}x")
# --- Test 5: vllm paged attention (decode) ---
print("\n--- Test 5: vllm_single_query_cached_kv_attention ---")
if hasattr(ixformer, 'vllm_single_query_cached_kv_attention'):
# Simulate decode with KV cache
# This is the function vllm uses for decode path
num_seqs = 1
num_kv_heads_total = num_kv_heads
block_size = 16
num_blocks = 64 # 64*16 = 1024 context tokens
max_context_len = num_blocks * block_size
q = torch.randn(num_seqs, num_heads, head_dim,
device="cuda", dtype=torch.float16)
k_cache = torch.randn(num_blocks * num_seqs, num_kv_heads_total,
head_dim, block_size,
device="cuda", dtype=torch.float16)
v_cache = torch.randn(num_blocks * num_seqs, num_kv_heads_total,
head_dim, block_size,
device="cuda", dtype=torch.float16)
block_tables = torch.arange(num_blocks, device="cuda",
dtype=torch.int32).unsqueeze(0)
context_lens = torch.tensor([max_context_len], device="cuda",
dtype=torch.int32)
scale = head_dim ** -0.5
out = torch.empty(num_seqs, num_heads, head_dim,
device="cuda", dtype=torch.float16)
try:
ixformer.vllm_single_query_cached_kv_attention(
out, q, k_cache, v_cache, scale,
block_tables, context_lens, block_size, max_context_len)
has_nan = out.isnan().any().item()
print(f" paged attn: shape={out.shape} nan={has_nan}")
except Exception as e:
print(f" paged attn: EXCEPTION: {e}")
print("\n" + "=" * 60)
return 0
if __name__ == "__main__":
sys.exit(main())