Download benchmark.py from Snapkitty/snapkitty-transformer: direct link, hf CLI and curl.
- Browser
- Download file 4.93 kB
-
https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/benchmark.py
- Command line
-
hf download hf://Snapkitty/snapkitty-transformer/benchmark.py
-
curl -L -o benchmark.py https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/benchmark.py
4.93 kB
| 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() |