File size: 4,002 Bytes
b3e26c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
"""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")