h3-blackwell-runtime/tools/smoke_flash4.py
2026-08-20 22:08:52 +07:00

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))