strix-halo-kernels / test_flash.py
axjns's picture
Add Triton flash attention (2-3x over AOTriton, 31x less peak memory)
b3e26c5 verified
Raw
History Blame Contribute Delete
4 kB
"""Correctness + headroom for the Triton flash attention on gfx1151.
Baseline is torch SDPA, which on this part can only use the `math` backend --
the thing we are trying to beat.
"""
import sys, torch, triton
import torch.nn.functional as F
from flash_attn import flash_attn
DEV = "cuda"
FAILS = []
print("gfx:", torch.cuda.get_device_properties(0).gcnArchName, "| triton", triton.__version__)
def ref(q, k, v, causal):
# fp32 reference so we compare against the true value, not the fp16 baseline
return F.scaled_dot_product_attention(q.float(), k.float(), v.float(), is_causal=causal)
def check(tag, got, want, atol):
d = (got.float() - want.float()).abs().max().item()
ok = d <= atol and torch.isfinite(got).all()
print(f" {'PASS' if ok else 'FAIL'} {tag:46} max|diff|={d:.3e}")
if not ok:
FAILS.append(tag)
print("\n== correctness vs fp32 reference ==")
for causal in (False, True):
for B, H, S, D in ((1, 2, 128, 64), (2, 4, 256, 64), (1, 8, 512, 128),
(1, 2, 1024, 128), (2, 2, 333, 64), (1, 1, 77, 64)):
for dt, atol in ((torch.float16, 4e-3), (torch.bfloat16, 3e-2)):
q = torch.randn(B, H, S, D, dtype=dt, device=DEV)
k = torch.randn(B, H, S, D, dtype=dt, device=DEV)
v = torch.randn(B, H, S, D, dtype=dt, device=DEV)
check(f"causal={causal} {dt.__str__().split('.')[-1]:9} ({B},{H},{S},{D})",
flash_attn(q, k, v, is_causal=causal), ref(q, k, v, causal), atol)
print("\n== guards ==")
for tag, fn in [
("shape mismatch", lambda: flash_attn(torch.randn(1,2,8,64,device=DEV),
torch.randn(1,2,9,64,device=DEV),
torch.randn(1,2,9,64,device=DEV))),
("3-D input", lambda: flash_attn(*[torch.randn(2,8,64,device=DEV)]*3)),
("non-pow2 head_dim", lambda: flash_attn(*[torch.randn(1,2,8,48,device=DEV)]*3)),
]:
try:
fn(); print(f" FAIL {tag} not rejected"); FAILS.append(tag)
except ValueError:
print(f" PASS {tag} rejected")
print("\n== headroom vs torch SDPA (math fallback), fp16, causal ==")
print(f" {'shape':>22} {'torch SDPA':>12} {'triton':>10} {'speedup':>9}")
rows = []
for B, H, S, D in ((1, 32, 512, 128), (1, 32, 4096, 128), (2, 24, 4096, 64), (1, 16, 8192, 64)):
q = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
k = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
v = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
a = triton.testing.do_bench(lambda: F.scaled_dot_product_attention(q, k, v, is_causal=True),
warmup=25, rep=100)
b = triton.testing.do_bench(lambda: flash_attn(q, k, v, is_causal=True), warmup=25, rep=100)
rows.append((f"({B},{H},{S},{D})", a, b))
print(f" {rows[-1][0]:>22} {a:10.3f}ms {b:8.3f}ms {a/b:8.2f}x")
print("\n== peak memory: materialized vs tiled ==")
for B, H, S, D in ((1, 32, 4096, 128),):
q = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
k = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
v = torch.randn(B, H, S, D, dtype=torch.float16, device=DEV)
for name, fn in (("torch SDPA", lambda: F.scaled_dot_product_attention(q, k, v, is_causal=True)),
("triton", lambda: flash_attn(q, k, v, is_causal=True))):
torch.cuda.synchronize(); torch.cuda.reset_peak_memory_stats()
fn(); torch.cuda.synchronize()
print(f" {name:12} peak alloc {torch.cuda.max_memory_allocated()/1024**2:9.1f} MiB")
best = _fwd_best = None
try:
from flash_attn import _fwd
k0 = next(iter(_fwd.cache.values())) if hasattr(_fwd, "cache") and _fwd.cache else None
if isinstance(_fwd.cache, dict) and _fwd.cache:
first = list(_fwd.cache.values())[0]
print("\nautotune picked:", first)
except Exception:
pass
print()
if FAILS:
print(f"{len(FAILS)} FAILURE(S): {FAILS}"); sys.exit(1)
print("all checks passed")