50 lines
1.8 KiB
Python
50 lines
1.8 KiB
Python
"""Compare FlashAttention-4 against PyTorch SDPA on an H3-shaped operation."""
|
|
|
|
import argparse
|
|
import json
|
|
|
|
import torch
|
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--sequence", type=int, default=257)
|
|
parser.add_argument("--heads", type=int, default=8)
|
|
parser.add_argument("--head-dim", type=int, default=128)
|
|
parser.add_argument("--iterations", type=int, default=5)
|
|
args = parser.parse_args()
|
|
|
|
from flash_attn.cute import flash_attn_func
|
|
|
|
torch.manual_seed(440411)
|
|
shape = (1, args.sequence, args.heads, args.head_dim)
|
|
q = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
|
|
k = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
|
|
v = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
|
|
|
|
with torch.inference_mode():
|
|
result = flash_attn_func(q, k, v, causal=False)
|
|
actual = result[0] if isinstance(result, tuple) else result
|
|
expected = torch.nn.functional.scaled_dot_product_attention(
|
|
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=False,
|
|
).transpose(1, 2)
|
|
torch.cuda.synchronize()
|
|
started = torch.cuda.Event(enable_timing=True)
|
|
finished = torch.cuda.Event(enable_timing=True)
|
|
started.record()
|
|
for _ in range(args.iterations):
|
|
result = flash_attn_func(q, k, v, causal=False)
|
|
actual = result[0] if isinstance(result, tuple) else result
|
|
finished.record()
|
|
torch.cuda.synchronize()
|
|
|
|
error = (actual.float() - expected.float()).abs()
|
|
print(json.dumps({
|
|
"device": torch.cuda.get_device_name(),
|
|
"capability": torch.cuda.get_device_capability(),
|
|
"shape": tuple(actual.shape),
|
|
"dtype": str(actual.dtype),
|
|
"contiguous": actual.is_contiguous(),
|
|
"max_abs_error": error.max().item(),
|
|
"mean_abs_error": error.mean().item(),
|
|
"milliseconds": started.elapsed_time(finished) / args.iterations,
|
|
}, indent=2))
|