import torch import time import snapkitty_flash_attention import json import csv from typing import List, Dict, Tuple class FlashAttentionBenchmark: def __init__(self, device='cuda', dtype=torch.half): self.device = device self.dtype = dtype self.results = [] def benchmark_config(self, batch, heads, seq_len, dim, num_warmup=10, num_iters=100, causal=False): Q = torch.randn(batch, heads, seq_len, dim, dtype=self.dtype, device=self.device) K = torch.randn(batch, heads, seq_len, dim, dtype=self.dtype, device=self.device) V = torch.randn(batch, heads, seq_len, dim, dtype=self.dtype, device=self.device) for _ in range(num_warmup): _ = snapkitty_flash_attention.flash_attention_fwd(Q, K, V) torch.cuda.synchronize() start = time.perf_counter() for _ in range(num_iters): O = snapkitty_flash_attention.flash_attention_fwd(Q, K, V) torch.cuda.synchronize() fwd_time = (time.perf_counter() - start) / num_iters * 1000 O = snapkitty_flash_attention.flash_attention_fwd(Q, K, V) loss = O.sum() for _ in range(num_warmup): loss.backward() Q.grad.zero_(); K.grad.zero_(); V.grad.zero_() torch.cuda.synchronize() start = time.perf_counter() for _ in range(num_iters): loss.backward() Q.grad.zero_(); K.grad.zero_(); V.grad.zero_() torch.cuda.synchronize() bwd_time = (time.perf_counter() - start) / num_iters * 1000 torch.cuda.reset_peak_memory_stats() _ = snapkitty_flash_attention.flash_attention_fwd(Q, K, V) torch.cuda.synchronize() peak_mem = torch.cuda.max_memory_allocated() / 1024**2 flops = 2 * batch * heads * seq_len * seq_len * dim + 2 * batch * heads * seq_len * dim tflops = flops / (fwd_time / 1000) / 1e12 result = { 'batch': batch, 'heads': heads, 'seq_len': seq_len, 'dim': dim, 'fwd_ms': fwd_time, 'bwd_ms': bwd_time, 'total_ms': fwd_time + bwd_time, 'peak_mem_mb': peak_mem, 'tflops': tflops, 'causal': causal } self.results.append(result) return result def run_sweep(self, configs): for config in configs: print(f"Benchmarking: B={config[0]}, H={config[1]}, N={config[2]}, D={config[3]}") try: result = self.benchmark_config(*config) print(f" FWD: {result['fwd_ms']:.2f} ms, BWD: {result['bwd_ms']:.2f} ms, TFLOPs: {result['tflops']:.2f}") except Exception as e: print(f" FAILED: {e}") return self.results def save_results(self, filepath): with open(filepath + '.json', 'w') as f: json.dump(self.results, f, indent=2) if self.results and 'error' not in str(self.results[0]): with open(filepath + '.csv', 'w', newline='') as f: writer = csv.DictWriter(f, fieldnames=self.results[0].keys()) writer.writeheader() writer.writerows(self.results) def print_summary(self): print("\n" + "="*100) print(f"{'Config':<30} {'FWD (ms)':>10} {'BWD (ms)':>10} {'Total (ms)':>12} {'TFLOPs':>8} {'Mem (MB)':>10}") print("-"*100) for r in self.results: cfg = f"B{r['batch']}H{r['heads']}N{r['seq_len']}D{r['dim']}" print(f"{cfg:<30} {r['fwd_ms']:>10.2f} {r['bwd_ms']:>10.2f} {r['total_ms']:>12.2f} {r['tflops']:>8.2f} {r['peak_mem_mb']:>10.0f}") print("="*100) def compare_with_pytorch(): print("\n=== Correctness vs PyTorch SDPA ===") torch.manual_seed(42) configs = [(2,8,512,64), (2,8,1024,64), (2,8,2048,64), (4,16,1024,64)] for batch, heads, seq_len, dim in configs: Q = torch.randn(batch, heads, seq_len, dim, dtype=torch.half, device='cuda') K = torch.randn(batch, heads, seq_len, dim, dtype=torch.half, device='cuda') V = torch.randn(batch, heads, seq_len, dim, dtype=torch.half, device='cuda') O_fa = snapkitty_flash_attention.flash_attention(Q, K, V) with torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=False, enable_mem_efficient=False): O_sdpa = torch.nn.functional.scaled_dot_product_attention(Q, K, V) max_diff = (O_fa - O_sdpa).abs().max().item() print(f"B={batch} H={heads} N={seq_len} D={dim}: max_diff={max_diff:.6f}") assert max_diff < 1e-3, f"Numerical mismatch: {max_diff}" print("All correctness tests passed!") if __name__ == '__main__': bench = FlashAttentionBenchmark() configs = [ (1,32,512,64), (1,32,1024,64), (1,32,2048,64), (1,32,4096,64), (2,16,1024,64), (4,8,1024,64), (4,32,1024,64), (8,32,1024,64), ] bench.run_sweep(configs) bench.print_summary() bench.save_results('benchmark_results') compare_with_pytorch()