//! Performance and benchmark tests for Flash Attention //! Validates 5-8x speedup claims and memory efficiency //! //! NOTE: Disabled until Flash Attention API is fully implemented #![cfg(all(feature = "cuda", feature = "disabled_tests"))] use criterion::{Criterion, black_box}; use rtx_flash_attention::*; use rtx_tensor::Tensor; use std::time::Instant; /// Memory usage tracker for tests struct MemoryTracker { peak_usage: usize, current_usage: usize, } impl MemoryTracker { fn new() -> Self { Self { peak_usage: 0, current_usage: 0, } } fn allocate(&mut self, size: usize) { self.current_usage += size; self.peak_usage = self.peak_usage.max(self.current_usage); } fn deallocate(&mut self, size: usize) { self.current_usage = self.current_usage.saturating_sub(size); } } /// Test that Flash Attention achieves target speedup on large sequences #[test] fn test_speedup_large_sequences() { let configs = vec![ (1024, 64, 8), // seq_len, head_dim, num_heads (2048, 64, 8), (4096, 64, 8), (8_192, 64, 4), // Reduced heads for very long sequences ]; for (seq_len, head_dim, num_heads) in configs { println!( "Testing seq_len={}, head_dim={}, num_heads={}", seq_len, head_dim, num_heads ); let config = FlashAttentionConfig { head_dim, block_size_q: 128, block_size_k: 128, causal: false, softmax_scale: None, }; let q = Tensor::randn(&[1, num_heads, seq_len, head_dim]); let k = Tensor::randn(&[1, num_heads, seq_len, head_dim]); let v = Tensor::randn(&[1, num_heads, seq_len, head_dim]); // Warmup let _ = flash_attention_forward(&q, &k, &v, &config).unwrap(); let _ = reference_attention(&q, &k, &v, config.softmax_scale).unwrap(); // Benchmark Flash Attention let start = Instant::now(); for _ in 0..10 { let _ = black_box(flash_attention_forward(&q, &k, &v, &config).unwrap()); } let flash_time = start.elapsed().as_millis() as f64 / 10.0; // Benchmark reference implementation let start = Instant::now(); for _ in 0..10 { let _ = black_box(reference_attention(&q, &k, &v, config.softmax_scale).unwrap()); } let reference_time = start.elapsed().as_millis() as f64 / 10.0; let speedup = reference_time / flash_time; println!( "Flash Attention: {:.2}ms, Reference: {:.2}ms, Speedup: {:.2}x", flash_time, reference_time, speedup ); // For long sequences, we should see significant speedup if seq_len >= 2048 { assert!( speedup >= 2.0, "Expected at least 2x speedup for seq_len={}, got {:.2}x", seq_len, speedup ); } if seq_len >= 4096 { assert!( speedup >= 3.0, "Expected at least 3x speedup for seq_len={}, got {:.2}x", seq_len, speedup ); } } } /// Test memory efficiency - O(sqrt(N)) vs O(N^2) #[test] fn test_memory_efficiency() { let sequence_lengths = vec![512, 1024, 2048, 4096]; for seq_len in sequence_lengths { let config = FlashAttentionConfig { head_dim: 64, block_size_q: 128, block_size_k: 128, causal: false, softmax_scale: None, }; let q = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let k = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let v = Tensor::randn(&[1, 1, seq_len, config.head_dim]); // Estimate memory usage for Flash Attention (should be O(sqrt(N))) let flash_memory = estimate_flash_attention_memory(&config, seq_len); // Estimate memory usage for standard attention (O(N^2)) let standard_memory = estimate_standard_attention_memory(seq_len, config.head_dim); let memory_ratio = standard_memory as f64 / flash_memory as f64; println!( "seq_len={}: Flash={}MB, Standard={}MB, Ratio={:.2}x", seq_len, flash_memory / 1024 / 1024, standard_memory / 1024 / 1024, memory_ratio ); // Memory savings should increase with sequence length if seq_len >= 1024 { assert!( memory_ratio >= 2.0, "Expected memory savings for seq_len={}", seq_len ); } if seq_len >= 2048 { assert!( memory_ratio >= 4.0, "Expected significant memory savings for seq_len={}", seq_len ); } } } /// Test throughput on different hardware configurations #[test] fn test_throughput_scaling() { let config = FlashAttentionConfig { head_dim: 64, block_size_q: 64, block_size_k: 64, causal: false, softmax_scale: None, }; let batch_sizes = vec![1, 2, 4, 8, 16]; let seq_len = 1024; for batch_size in batch_sizes { let q = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let k = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let v = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let start = Instant::now(); let _ = flash_attention_forward(&q, &k, &v, &config).unwrap(); let elapsed = start.elapsed().as_millis() as f64; let throughput = (batch_size as f64 * seq_len as f64) / elapsed * 1000.0; // tokens/second println!("Batch size {}: {:.0} tokens/second", batch_size, throughput); // Throughput should scale reasonably with batch size assert!(throughput > 0.0); if batch_size == 1 { // Ensure minimum throughput for single batch assert!( throughput > 1000.0, "Minimum throughput not met: {:.0}", throughput ); } } } /// Test numerical stability under extreme conditions #[test] fn test_numerical_stability() { let config = FlashAttentionConfig { head_dim: 64, block_size_q: 32, block_size_k: 32, causal: false, softmax_scale: Some(0.125), }; // Test with very large values let large_q = Tensor::full(&[1, 1, 128, config.head_dim], 10.0); let large_k = Tensor::full(&[1, 1, 128, config.head_dim], 10.0); let large_v = Tensor::full(&[1, 1, 128, config.head_dim], 10.0); let result = flash_attention_forward(&large_q, &large_k, &large_v, &config).unwrap(); assert!( result.all_finite().unwrap(), "Large values should remain stable" ); // Test with very small values let small_q = Tensor::full(&[1, 1, 128, config.head_dim], 1e-6); let small_k = Tensor::full(&[1, 1, 128, config.head_dim], 1e-6); let small_v = Tensor::full(&[1, 1, 128, config.head_dim], 1e-6); let result = flash_attention_forward(&small_q, &small_k, &small_v, &config).unwrap(); assert!( result.all_finite().unwrap(), "Small values should remain stable" ); // Test mixed precision scenarios let mixed_q = Tensor::randn(&[1, 1, 128, config.head_dim]); let mixed_k = mixed_q.clone() * 1000.0; // Very different scales let mixed_v = Tensor::full(&[1, 1, 128, config.head_dim], 0.001); let result = flash_attention_forward(&mixed_q, &mixed_k, &mixed_v, &config).unwrap(); assert!( result.all_finite().unwrap(), "Mixed precision should remain stable" ); } /// Test scalability across different sequence lengths #[test] fn test_scalability() { let config = FlashAttentionConfig { head_dim: 64, block_size_q: 128, block_size_k: 128, causal: false, softmax_scale: None, }; let mut previous_time = 0.0; let sequence_lengths = vec![256, 512, 1024, 2048, 4096]; for seq_len in sequence_lengths { let q = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let k = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let v = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let start = Instant::now(); let _ = flash_attention_forward(&q, &k, &v, &config).unwrap(); let elapsed = start.elapsed().as_millis() as f64; println!("seq_len={}: {:.2}ms", seq_len, elapsed); if previous_time > 0.0 { let time_ratio = elapsed / previous_time; let length_ratio = seq_len as f64 / (seq_len / 2) as f64; // 2x length increase // Time should scale better than O(N^2) due to Flash Attention optimization // For 2x sequence length, time should be much less than 4x assert!( time_ratio < 3.0, "Time scaling too poorly: {}x time for {}x length", time_ratio, length_ratio ); } previous_time = elapsed; } } /// Test concurrent execution performance #[test] fn test_concurrent_performance() { use std::sync::Arc; use std::thread; let config = Arc::new(FlashAttentionConfig { head_dim: 64, block_size_q: 64, block_size_k: 64, causal: false, softmax_scale: None, }); let seq_len = 512; let num_threads = 4; let iterations_per_thread = 10; let handles: Vec<_> = (0..num_threads) .map(|thread_id| { let config = Arc::clone(&config); thread::spawn(move || { let q = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let k = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let v = Tensor::randn(&[1, 1, seq_len, config.head_dim]); let start = Instant::now(); for _ in 0..iterations_per_thread { let _ = flash_attention_forward(&q, &k, &v, &config).unwrap(); } let elapsed = start.elapsed(); println!( "Thread {}: {:.2}ms per iteration", thread_id, elapsed.as_millis() as f64 / iterations_per_thread as f64 ); elapsed }) }) .collect(); let results: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect(); // Verify all threads completed successfully assert_eq!(results.len(), num_threads); // Check for reasonable performance consistency across threads let times: Vec = results.iter().map(|d| d.as_millis() as f64).collect(); let mean_time = times.iter().sum::() / times.len() as f64; for &time in × { let deviation = (time - mean_time).abs() / mean_time; assert!( deviation < 0.5, "Thread performance variance too high: {:.1}%", deviation * 100.0 ); } } /// Test gradient computation performance #[test] fn test_gradient_performance() { let config = FlashAttentionConfig { head_dim: 64, block_size_q: 64, block_size_k: 64, causal: false, softmax_scale: None, }; let seq_len = 1024; let q = Tensor::randn(&[1, 8, seq_len, config.head_dim]); let k = Tensor::randn(&[1, 8, seq_len, config.head_dim]); let v = Tensor::randn(&[1, 8, seq_len, config.head_dim]); // Forward pass timing let start = Instant::now(); let output = flash_attention_forward(&q, &k, &v, &config).unwrap(); let forward_time = start.elapsed().as_millis() as f64; // Backward pass timing let grad_output = Tensor::randn_like(&output); let start = Instant::now(); let _ = flash_attention_backward(&grad_output, &q, &k, &v, &config).unwrap(); let backward_time = start.elapsed().as_millis() as f64; println!( "Forward: {:.2}ms, Backward: {:.2}ms", forward_time, backward_time ); // Backward pass should be at most 3x forward pass time assert!( backward_time < forward_time * 3.0, "Backward pass too slow: {:.2}ms vs {:.2}ms forward", backward_time, forward_time ); } /// Test block size optimization #[test] fn test_block_size_optimization() { let seq_len = 1024; let head_dim = 64; let q = Tensor::randn(&[1, 1, seq_len, head_dim]); let k = Tensor::randn(&[1, 1, seq_len, head_dim]); let v = Tensor::randn(&[1, 1, seq_len, head_dim]); let block_sizes = vec![32, 64, 128, 256]; let mut best_time = f64::INFINITY; let mut best_block_size = 0; for block_size in block_sizes { if block_size <= seq_len { let config = FlashAttentionConfig { head_dim, block_size_q: block_size, block_size_k: block_size, causal: false, softmax_scale: None, }; let start = Instant::now(); for _ in 0..5 { let _ = flash_attention_forward(&q, &k, &v, &config).unwrap(); } let elapsed = start.elapsed().as_millis() as f64 / 5.0; println!("Block size {}: {:.2}ms", block_size, elapsed); if elapsed < best_time { best_time = elapsed; best_block_size = block_size; } } } println!( "Optimal block size: {} ({:.2}ms)", best_block_size, best_time ); // Verify we found a reasonable block size assert!(best_block_size >= 32 && best_block_size <= 256); assert!(best_time < f64::INFINITY); } /// Helper functions for memory estimation fn estimate_flash_attention_memory(config: &FlashAttentionConfig, seq_len: usize) -> usize { // Flash Attention memory is roughly O(sqrt(N)) due to block-wise computation let block_memory = config.block_size_q * config.block_size_k * 4; // fp32 let intermediate_memory = config.block_size_q * config.head_dim * 4; let total_blocks = (seq_len + config.block_size_q - 1) / config.block_size_q; block_memory + intermediate_memory + (total_blocks * 64) // Some overhead per block } fn estimate_standard_attention_memory(seq_len: usize, head_dim: usize) -> usize { // Standard attention memory is O(N^2) for attention matrix let attention_matrix = seq_len * seq_len * 4; // fp32 let intermediate = seq_len * head_dim * 4 * 3; // Q, K, V copies attention_matrix + intermediate } /// Reference implementation for performance comparison fn reference_attention( q: &Tensor, k: &Tensor, v: &Tensor, scale: Option, ) -> Result> { let scale = scale.unwrap_or(1.0 / (q.shape()[3] as f32).sqrt()); // QK^T - this creates the O(N^2) memory bottleneck let scores = q.matmul(&k.transpose(-2, -1)?)?; let scaled_scores = scores.mul_scalar(scale)?; // Softmax - also O(N^2) memory let attention_probs = scaled_scores.softmax(-1)?; // Apply to values let output = attention_probs.matmul(v)?; Ok(output) } /// Test load balancing across multiple GPUs (if available) #[test] fn test_multi_gpu_load_balancing() { // This test would be more meaningful with actual multi-GPU setup // For now, test that we can handle different batch sizes efficiently let config = FlashAttentionConfig { head_dim: 64, block_size_q: 64, block_size_k: 64, causal: false, softmax_scale: None, }; let batch_sizes = vec![1, 4, 8, 16]; let seq_len = 512; for batch_size in batch_sizes { let q = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let k = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let v = Tensor::randn(&[batch_size, 8, seq_len, config.head_dim]); let start = Instant::now(); let result = flash_attention_forward(&q, &k, &v, &config).unwrap(); let elapsed = start.elapsed().as_millis() as f64; let throughput_per_sample = elapsed / batch_size as f64; println!( "Batch {}: {:.2}ms total, {:.2}ms per sample", batch_size, elapsed, throughput_per_sample ); // Verify results are correct shape assert_eq!(result.shape(), &[batch_size, 8, seq_len, config.head_dim]); // Larger batches should have better throughput per sample (GPU utilization) if batch_size >= 4 { assert!( throughput_per_sample < 50.0, "Poor GPU utilization: {:.2}ms per sample", throughput_per_sample ); } } }