snapkitty-transformer / benchmark.py
SNAPKITTYWEST's picture
chore: push from SNAPKITTYWEST local build
fd6abd3 verified
Raw History Blame Contribute Delete
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()