//! Flash Attention Benchmark - Metal vs CPU with Causal Masking //! Run with: cargo run -p rtx-flash-metal-attention --example simple_bench --release use rtx_flash_metal_attention::{FlashAttention, FlashAttentionConfig}; use rtx_tensor::{Device, Tensor}; use std::time::Instant; /// Naive CPU attention implementation for comparison /// O = softmax(Q @ K^T / sqrt(d)) @ V fn cpu_naive_attention( batch_size: usize, num_heads: usize, seq_len: usize, head_dim: usize, causal: bool, ) -> f64 { let scale = 1.0 / (head_dim as f32).sqrt(); // Simulate Q, K, V as all ones (same as Metal benchmark) let q = vec![1.0f32; batch_size * num_heads * seq_len * head_dim]; let k = vec![1.0f32; batch_size * num_heads * seq_len * head_dim]; let v = vec![1.0f32; batch_size * num_heads * seq_len * head_dim]; let mut output = vec![0.0f32; batch_size * num_heads * seq_len * head_dim]; // For each batch and head for b in 0..batch_size { for h in 0..num_heads { let base = (b * num_heads + h) * seq_len * head_dim; // For each query position for i in 0..seq_len { // Compute attention scores for row i let mut scores = vec![f32::NEG_INFINITY; seq_len]; let mut max_score = f32::NEG_INFINITY; // Apply causal mask: only attend to positions <= i let end_j = if causal { i + 1 } else { seq_len }; for j in 0..end_j { let mut dot = 0.0f32; for d in 0..head_dim { dot += q[base + i * head_dim + d] * k[base + j * head_dim + d]; } scores[j] = dot * scale; max_score = max_score.max(scores[j]); } // Softmax (only over valid positions) let mut sum_exp = 0.0f32; for j in 0..end_j { scores[j] = (scores[j] - max_score).exp(); sum_exp += scores[j]; } for j in 0..end_j { scores[j] /= sum_exp; } // Weighted sum of V for d in 0..head_dim { let mut acc = 0.0f32; for j in 0..end_j { acc += scores[j] * v[base + j * head_dim + d]; } output[base + i * head_dim + d] = acc; } } } } // Return sum to prevent optimization output.iter().map(|x| *x as f64).sum() } fn bench_attention( metal_device: &Device, config: FlashAttentionConfig, scenarios: &[(&str, usize, usize, usize)], causal: bool, ) -> Vec { let mode = if causal { "Causal" } else { "Full" }; let attn = match FlashAttention::new(config) { Ok(a) => a, Err(e) => { eprintln!("Failed to create FlashAttention: {:?}", e); return vec![]; } }; println!("\n{:=^80}", format!(" Metal Flash Attention ({}) ", mode)); println!( "{:<20} {:>15} {:>15} {:>15}", "Scenario", "Avg Time (ms)", "Throughput", "Elements" ); println!("{:-<80}", ""); let mut metal_times = Vec::new(); for (name, batch_size, seq_len, head_dim) in scenarios { let num_heads = 8; let shape = &[*batch_size, num_heads, *seq_len, *head_dim]; let q = Tensor::ones(shape, metal_device).expect("Q tensor"); let k = Tensor::ones(shape, metal_device).expect("K tensor"); let v = Tensor::ones(shape, metal_device).expect("V tensor"); // Warmup for _ in 0..10 { let _ = attn.forward(&q, &k, &v).expect("warmup forward"); } // Benchmark let iterations = 100; let start = Instant::now(); for _ in 0..iterations { let _ = attn.forward(&q, &k, &v).expect("forward"); } let elapsed = start.elapsed(); let avg_ms = elapsed.as_secs_f64() * 1000.0 / iterations as f64; let total_elements = (batch_size * num_heads * seq_len * head_dim) as f64; let throughput = total_elements / (avg_ms / 1000.0) / 1e6; metal_times.push(avg_ms); println!( "{:<20} {:>12.4} ms {:>12.2} M/s {:>12.0}", name, avg_ms, throughput, total_elements ); } println!("\n{:=^80}", format!(" CPU Naive Attention ({}) ", mode)); println!( "{:<20} {:>15} {:>15} {:>15}", "Scenario", "Avg Time (ms)", "Throughput", "Speedup" ); println!("{:-<80}", ""); for (idx, (name, batch_size, seq_len, head_dim)) in scenarios.iter().enumerate() { let num_heads = 8; let total_elements = (batch_size * num_heads * seq_len * head_dim) as f64; // Warmup for _ in 0..2 { let _ = cpu_naive_attention(*batch_size, num_heads, *seq_len, *head_dim, causal); } // Fewer iterations for CPU (it's slower) let iterations = 10; let start = Instant::now(); for _ in 0..iterations { let _ = cpu_naive_attention(*batch_size, num_heads, *seq_len, *head_dim, causal); } let elapsed = start.elapsed(); let avg_ms = elapsed.as_secs_f64() * 1000.0 / iterations as f64; let throughput = total_elements / (avg_ms / 1000.0) / 1e6; let speedup = avg_ms / metal_times[idx]; println!( "{:<20} {:>12.4} ms {:>12.2} M/s {:>12.1}x", name, avg_ms, throughput, speedup ); } metal_times } fn main() { println!("=== Flash Attention: Metal vs CPU Benchmark ===\n"); // Get Metal device let metal_device = match Device::metal(0) { Ok(d) => { println!("Metal device: Metal(0)"); d } Err(e) => { eprintln!("Metal device not available: {:?}", e); return; } }; // Test scenarios (head_dim must be <= max_head_dim which is 64) let scenarios = vec![ ("Latency_BS1", 1usize, 128usize, 64usize), ("Throughput_BS32", 32, 128, 64), ("Heavy_BS128", 128, 128, 64), ]; // Benchmark full attention (non-causal) println!("\n{:#^80}", " FULL ATTENTION (Non-Causal) "); let config_full = FlashAttentionConfig::default(); let _full_times = bench_attention(&metal_device, config_full, &scenarios, false); // Benchmark causal attention println!("\n\n{:#^80}", " CAUSAL ATTENTION (Autoregressive) "); let config_causal = FlashAttentionConfig::causal(); let _causal_times = bench_attention(&metal_device, config_causal, &scenarios, true); // Benchmark with larger head dimension (128) println!("\n\n{:#^80}", " LARGE HEAD DIM (128) "); let large_scenarios = vec![ ("BS1_HD128", 1usize, 128usize, 128usize), ("BS32_HD128", 32, 128, 128), ("BS64_HD128", 64, 128, 128), ]; let config_large = FlashAttentionConfig::large_head_dim(); let _large_times = bench_attention(&metal_device, config_large, &large_scenarios, false); println!("\n{:-<80}", ""); println!("\nDone!"); }