Files
rustytorch/crates/training/rtx-flash-attention/tests/performance_tests.rs
T
2026-03-04 00:08:42 +00:00

523 lines
16 KiB
Rust

//! 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<f64> = results.iter().map(|d| d.as_millis() as f64).collect();
let mean_time = times.iter().sum::<f64>() / times.len() as f64;
for &time in &times {
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<f32>,
) -> Result<Tensor, Box<dyn std::error::Error>> {
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
);
}
}
}