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