#!/usr/bin/env python3 """ PyTorch MPS Flash Attention Benchmark - Comparison baseline for RustyTorch Metal. This script benchmarks PyTorch's scaled_dot_product_attention on Apple MPS backend and outputs JSON results for automated comparison with RustyTorch Metal Flash Attention. Usage: python3 benchmarks/metal/bench_flash_attention.py python3 benchmarks/metal/bench_flash_attention.py --json > benchmarks/reports/pytorch_flash.json Run Rust benchmark with: cargo bench -p rtx-flash-metal-attention --bench metal_vs_pytorch """ import torch import torch.nn.functional as F import time import sys import json import argparse import platform import statistics from dataclasses import dataclass, asdict from typing import List, Optional @dataclass class BenchmarkResult: """Single benchmark result.""" name: str batch_size: int seq_len: int num_heads: int head_dim: int causal: bool avg_time_ms: float std_time_ms: float min_time_ms: float max_time_ms: float p50_time_ms: float p95_time_ms: float p99_time_ms: float throughput_elements_per_sec: float memory_mb: Optional[float] = None @dataclass class BenchmarkReport: """Complete benchmark report.""" framework: str backend: str pytorch_version: str python_version: str macos_version: str chip: str timestamp: str iterations: int warmup: int results: List[dict] def check_mps() -> bool: """Check if MPS backend is available.""" if not torch.backends.mps.is_available(): print("ERROR: MPS not available on this system", file=sys.stderr) print("This benchmark requires macOS with Apple Silicon", file=sys.stderr) return False return True def benchmark_attention( batch_size: int, seq_len: int = 128, head_dim: int = 64, num_heads: int = 8, causal: bool = False, iterations: int = 1000, warmup: int = 50, dtype=torch.float32, ) -> BenchmarkResult: """ Run a single benchmark scenario. Args: batch_size: Number of sequences in batch seq_len: Sequence length head_dim: Dimension per attention head num_heads: Number of attention heads causal: Whether to use causal masking iterations: Number of timed iterations warmup: Number of warmup iterations dtype: Data type (float32 to match Rust implementation) Returns: BenchmarkResult with timing statistics """ device = torch.device("mps") # Create input tensors [batch, heads, seq, head_dim] q = torch.randn(batch_size, num_heads, seq_len, head_dim, device=device, dtype=dtype) k = torch.randn(batch_size, num_heads, seq_len, head_dim, device=device, dtype=dtype) v = torch.randn(batch_size, num_heads, seq_len, head_dim, device=device, dtype=dtype) # Warmup: wake up GPU and JIT compile kernels for _ in range(warmup): _ = F.scaled_dot_product_attention(q, k, v, is_causal=causal) torch.mps.synchronize() # Collect timing samples times_ms = [] for _ in range(iterations): start = time.perf_counter() _ = F.scaled_dot_product_attention(q, k, v, is_causal=causal) torch.mps.synchronize() # Block until GPU done end = time.perf_counter() times_ms.append((end - start) * 1000) # Calculate statistics times_sorted = sorted(times_ms) avg_time = statistics.mean(times_ms) std_time = statistics.stdev(times_ms) if len(times_ms) > 1 else 0.0 min_time = times_sorted[0] max_time = times_sorted[-1] p50_idx = int(len(times_sorted) * 0.50) p95_idx = int(len(times_sorted) * 0.95) p99_idx = int(len(times_sorted) * 0.99) p50_time = times_sorted[p50_idx] p95_time = times_sorted[p95_idx] p99_time = times_sorted[min(p99_idx, len(times_sorted) - 1)] # Calculate throughput total_elements = batch_size * num_heads * seq_len * head_dim throughput = total_elements / (avg_time / 1000) # Get memory usage memory_mb = None try: # MPS doesn't have the same memory API as CUDA, estimate from tensor sizes tensor_bytes = q.numel() * q.element_size() * 3 # Q, K, V memory_mb = tensor_bytes / (1024 * 1024) except Exception: pass name = f"{'Causal_' if causal else ''}BS{batch_size}_Seq{seq_len}" return BenchmarkResult( name=name, batch_size=batch_size, seq_len=seq_len, num_heads=num_heads, head_dim=head_dim, causal=causal, avg_time_ms=avg_time, std_time_ms=std_time, min_time_ms=min_time, max_time_ms=max_time, p50_time_ms=p50_time, p95_time_ms=p95_time, p99_time_ms=p99_time, throughput_elements_per_sec=throughput, memory_mb=memory_mb, ) def run_benchmark_suite(iterations: int = 1000, warmup: int = 50) -> BenchmarkReport: """Run the complete benchmark suite.""" import datetime results = [] # Standard attention benchmarks (matching Rust scenarios) scenarios = [ # (batch_size, seq_len, head_dim, num_heads, causal) (1, 128, 64, 8, False), # Latency_BS1 - Inference (32, 128, 64, 8, False), # Throughput_BS32 (64, 128, 64, 8, False), # Throughput_BS64 - Training (256, 128, 64, 8, False), # Heavy_BS256 - GPU Saturation (1, 128, 64, 8, True), # Causal_BS1 - Autoregressive (32, 256, 64, 8, True), # Causal_BS32_Seq256 (64, 512, 64, 8, False), # Long_Seq512 (16, 1024, 64, 8, False), # Very_Long_Seq1024 ] for batch_size, seq_len, head_dim, num_heads, causal in scenarios: try: result = benchmark_attention( batch_size=batch_size, seq_len=seq_len, head_dim=head_dim, num_heads=num_heads, causal=causal, iterations=iterations, warmup=warmup, ) results.append(asdict(result)) except Exception as e: print(f"Warning: Benchmark failed for BS{batch_size}: {e}", file=sys.stderr) return BenchmarkReport( framework="PyTorch", backend="MPS", pytorch_version=torch.__version__, python_version=platform.python_version(), macos_version=platform.mac_ver()[0], chip=platform.processor() or "Apple Silicon", timestamp=datetime.datetime.now().isoformat(), iterations=iterations, warmup=warmup, results=results, ) def print_table(report: BenchmarkReport): """Print results as formatted table.""" print("=" * 90) print("PyTorch MPS Flash Attention Benchmark") print("=" * 90) print(f"PyTorch version: {report.pytorch_version}") print(f"Backend: {report.backend}") print(f"macOS version: {report.macos_version}") print(f"Chip: {report.chip}") print(f"Iterations: {report.iterations}") print() print("-" * 90) print(f"| {'Scenario':<25} | {'Avg (ms)':>10} | {'P50 (ms)':>10} | {'P99 (ms)':>10} | {'Throughput':>15} |") print("-" * 90) for r in report.results: throughput_str = f"{r['throughput_elements_per_sec']/1e6:.2f} M/s" print(f"| {r['name']:<25} | {r['avg_time_ms']:>10.4f} | {r['p50_time_ms']:>10.4f} | {r['p99_time_ms']:>10.4f} | {throughput_str:>15} |") print("-" * 90) print() print("To compare with RustyTorch Metal Flash Attention:") print(" cargo bench -p rtx-flash-metal-attention --bench metal_vs_pytorch") def main(): parser = argparse.ArgumentParser(description="PyTorch MPS Flash Attention Benchmark") parser.add_argument("--json", action="store_true", help="Output JSON format") parser.add_argument("--iterations", type=int, default=1000, help="Number of iterations") parser.add_argument("--warmup", type=int, default=50, help="Number of warmup iterations") args = parser.parse_args() if not check_mps(): sys.exit(1) report = run_benchmark_suite(iterations=args.iterations, warmup=args.warmup) if args.json: print(json.dumps(asdict(report), indent=2)) else: print_table(report) if __name__ == "__main__": main()