//! Configuration types for Flash Attention use crate::error::{FlashError, FlashResult}; use serde::{Deserialize, Serialize}; /// Flash Attention configuration #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct FlashAttentionConfig { /// Number of attention heads pub num_heads: usize, /// Dimension of each head pub head_dim: usize, /// Query block size for SRAM tiling (must be multiple of 32) pub block_size_q: usize, /// Key-Value block size for SRAM tiling (must be multiple of 32) pub block_size_kv: usize, /// Whether to use causal masking pub causal: bool, /// Softmax scaling factor (default: 1.0 / sqrt(head_dim)) pub softmax_scale: Option, /// Maximum sequence length supported pub max_seq_len: usize, /// CUDA device ID to use pub device_id: i32, /// Memory optimization level pub memory_optimization: MemoryOptimization, /// Numerical precision mode pub precision: PrecisionMode, /// Backend-specific configurations pub backend_config: BackendConfig, } impl Default for FlashAttentionConfig { fn default() -> Self { Self::new(8, 64) // Sensible defaults: 8 heads, 64 head dimension } } /// Memory optimization strategies #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum MemoryOptimization { /// Maximum speed, highest memory usage Speed, /// Balanced speed and memory Balanced, /// Minimum memory, may sacrifice some speed Memory, /// Custom optimization with specific parameters Custom { block_size_q: usize, block_size_kv: usize, sram_fraction: f32, }, } /// Numerical precision modes #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum PrecisionMode { /// Full 32-bit floating point FP32, /// Half precision (16-bit) FP16, /// Brain floating point (16-bit) BF16, /// FP8 E4M3 format - optimal for inference (better precision) FP8E4M3 { /// Configuration for FP8 quantization config: FP8Config, }, /// FP8 E5M2 format - optimal for training gradients (larger range) FP8E5M2 { /// Configuration for FP8 quantization config: FP8Config, }, /// Mixed precision with automatic scaling Mixed { compute_precision: Precision, storage_precision: Precision, }, } /// FP8 quantization configuration #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct FP8Config { /// Scaling factor for FP8 conversion /// If None, dynamic per-tensor scaling is used pub scale: Option, /// Use per-tensor dynamic scaling (recommended for accuracy) pub use_dynamic_scaling: bool, /// Compute precision for intermediate calculations /// FP32 recommended for numerical stability in softmax pub compute_precision: Precision, /// Gradient precision for backward pass /// Should be at least FP16 for training stability pub gradient_precision: Precision, /// Amax history length for dynamic scaling /// Longer history = more stable scaling, but slower adaptation pub amax_history_len: usize, /// Enable delayed scaling (update scale every N iterations) /// Reduces overhead but may impact accuracy pub delayed_scaling: bool, /// Number of iterations between scale updates (if delayed_scaling is true) pub scale_update_interval: usize, } impl Default for FP8Config { fn default() -> Self { Self { scale: None, use_dynamic_scaling: true, compute_precision: Precision::FP32, // Softmax needs FP32 for stability gradient_precision: Precision::BF16, // BF16 for gradients amax_history_len: 1024, delayed_scaling: false, scale_update_interval: 1, } } } impl FP8Config { /// Create FP8 config optimized for inference pub fn for_inference() -> Self { Self { scale: None, use_dynamic_scaling: true, compute_precision: Precision::FP16, // FP16 compute is faster gradient_precision: Precision::FP16, // Not used in inference amax_history_len: 256, delayed_scaling: true, scale_update_interval: 10, } } /// Create FP8 config optimized for training pub fn for_training() -> Self { Self { scale: None, use_dynamic_scaling: true, compute_precision: Precision::FP32, // FP32 for numerical stability gradient_precision: Precision::BF16, // BF16 for gradients amax_history_len: 1024, delayed_scaling: false, scale_update_interval: 1, } } /// Create FP8 config with static scaling factor pub fn with_static_scale(scale: f32) -> Self { Self { scale: Some(scale), use_dynamic_scaling: false, compute_precision: Precision::FP32, gradient_precision: Precision::BF16, amax_history_len: 0, delayed_scaling: false, scale_update_interval: 1, } } /// Validate FP8 configuration pub fn validate(&self) -> FlashResult<()> { if let Some(scale) = self.scale { if scale <= 0.0 || !scale.is_finite() { return Err(FlashError::config("FP8 scale must be positive and finite")); } } if self.delayed_scaling && self.scale_update_interval == 0 { return Err(FlashError::config("scale_update_interval must be > 0 when delayed_scaling is enabled")); } // Ensure compute precision is at least FP16 for softmax stability if matches!(self.compute_precision, Precision::INT8 | Precision::FP8E4M3 | Precision::FP8E5M2) { return Err(FlashError::config("FP8 compute_precision must be FP16 or higher for numerical stability")); } Ok(()) } } /// Precision types #[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] pub enum Precision { /// 32-bit floating point FP32, /// 16-bit half precision FP16, /// 16-bit brain floating point BF16, /// 8-bit integer (legacy, prefer FP8) INT8, /// 8-bit floating point E4M3 format (4 exponent, 3 mantissa) /// Better precision, smaller dynamic range /// Ideal for weights and activations in inference FP8E4M3, /// 8-bit floating point E5M2 format (5 exponent, 2 mantissa) /// Larger dynamic range, less precision /// Ideal for gradients in training FP8E5M2, } impl Precision { /// Get the size in bytes for this precision pub fn size_bytes(&self) -> usize { match self { Precision::FP32 => 4, Precision::FP16 | Precision::BF16 => 2, Precision::INT8 | Precision::FP8E4M3 | Precision::FP8E5M2 => 1, } } /// Check if this is an FP8 format pub fn is_fp8(&self) -> bool { matches!(self, Precision::FP8E4M3 | Precision::FP8E5M2) } /// Get the minimum compute capability required for this precision pub fn min_compute_capability(&self) -> (u32, u32) { match self { Precision::FP32 => (3, 0), Precision::FP16 => (5, 3), Precision::BF16 => (8, 0), Precision::INT8 => (6, 1), Precision::FP8E4M3 | Precision::FP8E5M2 => (8, 9), // Hopper (H100) or Ada (RTX 40xx) } } } /// Backend-specific configurations #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[derive(Default)] pub struct BackendConfig { /// Standard CUDA backend configuration pub cuda: CudaConfig, /// Edge deployment configuration #[cfg(feature = "edge")] pub edge: Option, } /// CUDA backend configuration #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct CudaConfig { /// CUDA stream pool size pub stream_pool_size: usize, /// Enable CUDA graphs for optimization pub enable_cuda_graphs: bool, /// Kernel auto-tuning parameters pub auto_tune: bool, /// Memory pool configuration pub memory_pool: MemoryPoolConfig, /// Compute capability target pub compute_capability: (u32, u32), } /// Memory pool configuration #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct MemoryPoolConfig { /// Initial pool size in bytes pub initial_size: usize, /// Maximum pool size in bytes pub max_size: usize, /// Memory growth factor pub growth_factor: f32, /// Enable memory defragmentation pub enable_defragmentation: bool, } /// Edge deployment configuration #[cfg(feature = "edge")] #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct EdgeConfig { /// Target platform pub target: EdgeTarget, /// Memory constraints pub memory_limit_mb: usize, /// Power constraints pub power_limit_mw: Option, /// Latency requirements pub max_latency_ms: Option, /// Quantization settings pub quantization: EdgeQuantization, } #[cfg(feature = "edge")] #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum EdgeTarget { RISCV, ARM { cortex: String }, WASM, Custom { architecture: String }, } #[cfg(feature = "edge")] #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct EdgeQuantization { pub weights: Precision, pub activations: Precision, pub enable_dynamic: bool, } impl FlashAttentionConfig { /// Create a new Flash Attention configuration with sensible defaults pub fn new(num_heads: usize, head_dim: usize) -> Self { let softmax_scale = 1.0 / (head_dim as f32).sqrt(); Self { num_heads, head_dim, block_size_q: 64, // Optimized for most GPUs block_size_kv: 64, causal: false, softmax_scale: Some(softmax_scale), max_seq_len: 32_768, // Support up to 32K context device_id: 0, memory_optimization: MemoryOptimization::Balanced, precision: PrecisionMode::FP16, backend_config: BackendConfig::default(), } } /// Create configuration optimized for training pub fn for_training(num_heads: usize, head_dim: usize) -> Self { let mut config = Self::new(num_heads, head_dim); config.memory_optimization = MemoryOptimization::Speed; config.precision = PrecisionMode::Mixed { compute_precision: Precision::FP32, storage_precision: Precision::FP16, }; config.backend_config.cuda.enable_cuda_graphs = true; config } /// Create configuration optimized for inference pub fn for_inference(num_heads: usize, head_dim: usize) -> Self { let mut config = Self::new(num_heads, head_dim); config.memory_optimization = MemoryOptimization::Memory; config.precision = PrecisionMode::FP16; config.backend_config.cuda.memory_pool.enable_defragmentation = true; config } /// Create configuration optimized for FP8 inference (2x memory reduction) /// Requires Hopper (H100) or Ada (RTX 40xx) GPUs with compute capability 8.9+ pub fn for_fp8_inference(num_heads: usize, head_dim: usize) -> Self { let mut config = Self::new(num_heads, head_dim); config.memory_optimization = MemoryOptimization::Memory; config.precision = PrecisionMode::FP8E4M3 { config: FP8Config::for_inference(), }; config.backend_config.cuda.memory_pool.enable_defragmentation = true; // Require Ada/Hopper for FP8 support config.backend_config.cuda.compute_capability = (8, 9); config } /// Create configuration optimized for FP8 training /// Uses FP8 for forward pass, higher precision for gradients pub fn for_fp8_training(num_heads: usize, head_dim: usize) -> Self { let mut config = Self::new(num_heads, head_dim); config.memory_optimization = MemoryOptimization::Balanced; // Use E4M3 for forward (better precision) with training config config.precision = PrecisionMode::FP8E4M3 { config: FP8Config::for_training(), }; config.backend_config.cuda.enable_cuda_graphs = true; config.backend_config.cuda.compute_capability = (8, 9); config } /// Validate the configuration pub fn validate(&self) -> FlashResult<()> { // Check head dimensions if self.num_heads == 0 { return Err(FlashError::config("num_heads must be greater than 0")); } if self.head_dim == 0 { return Err(FlashError::config("head_dim must be greater than 0")); } if !self.head_dim.is_multiple_of(8) { return Err(FlashError::config("head_dim must be multiple of 8 for vectorization")); } // Check block sizes if !self.block_size_q.is_multiple_of(32) { return Err(FlashError::config("block_size_q must be multiple of 32")); } if !self.block_size_kv.is_multiple_of(32) { return Err(FlashError::config("block_size_kv must be multiple of 32")); } if self.block_size_q > 1024 { return Err(FlashError::config("block_size_q too large (max 1024)")); } if self.block_size_kv > 1024 { return Err(FlashError::config("block_size_kv too large (max 1024)")); } // Check sequence length if self.max_seq_len == 0 { return Err(FlashError::config("max_seq_len must be greater than 0")); } // Check softmax scale if let Some(scale) = self.softmax_scale && (scale <= 0.0 || !scale.is_finite()) { return Err(FlashError::config("softmax_scale must be positive and finite")); } // Check device ID if self.device_id < 0 { return Err(FlashError::config("device_id must be non-negative")); } // Validate FP8 configuration if present match &self.precision { PrecisionMode::FP8E4M3 { config } | PrecisionMode::FP8E5M2 { config } => { config.validate()?; // Check compute capability requirements for FP8 let (major, minor) = self.backend_config.cuda.compute_capability; let required = Precision::FP8E4M3.min_compute_capability(); if major < required.0 || (major == required.0 && minor < required.1) { return Err(FlashError::config( format!( "FP8 requires compute capability {}.{} or higher (have {}.{}). \ FP8 is supported on H100 (Hopper) and RTX 40xx (Ada) GPUs.", required.0, required.1, major, minor ) )); } } _ => {} } self.backend_config.validate()?; Ok(()) } /// Get effective softmax scale pub fn get_softmax_scale(&self) -> f32 { self.softmax_scale.unwrap_or_else(|| 1.0 / (self.head_dim as f32).sqrt()) } /// Get optimal block sizes for the current configuration pub fn get_optimal_block_sizes(&self) -> (usize, usize) { match &self.memory_optimization { MemoryOptimization::Speed => (128, 128), MemoryOptimization::Balanced => (64, 64), MemoryOptimization::Memory => (32, 32), MemoryOptimization::Custom { block_size_q, block_size_kv, .. } => { (*block_size_q, *block_size_kv) } } } /// Calculate memory requirements in bytes pub fn estimate_memory_usage(&self, batch_size: usize, seq_len: usize) -> usize { let (storage_size, compute_size) = match &self.precision { PrecisionMode::FP32 => (4, 4), PrecisionMode::FP16 | PrecisionMode::BF16 => (2, 2), PrecisionMode::FP8E4M3 { config } | PrecisionMode::FP8E5M2 { config } => { // FP8 storage, but compute in higher precision let compute = config.compute_precision.size_bytes(); (1, compute) }, PrecisionMode::Mixed { storage_precision, compute_precision } => { (storage_precision.size_bytes(), compute_precision.size_bytes()) } }; // Memory for Q, K, V tensors (stored in storage precision) let qkv_memory = 3 * batch_size * self.num_heads * seq_len * self.head_dim * storage_size; // Memory for output tensor (stored in storage precision) let output_memory = batch_size * self.num_heads * seq_len * self.head_dim * storage_size; // Memory for intermediate computations (computed in compute precision) // Note: Flash Attention avoids materializing full attention matrix, // but we need workspace for tiled computation let intermediate_memory = batch_size * self.num_heads * self.block_size_q * self.block_size_kv * compute_size; // SRAM working memory for tiles (in compute precision) let sram_memory = self.block_size_q * self.block_size_kv * compute_size; // FP8 scaling factors (one per tensor if dynamic scaling) let scaling_overhead = match &self.precision { PrecisionMode::FP8E4M3 { config } | PrecisionMode::FP8E5M2 { config } if config.use_dynamic_scaling => { // Scale factors for Q, K, V, and output (4 tensors * 4 bytes per scale) 4 * 4 + config.amax_history_len * 4 // amax history buffer }, _ => 0, }; qkv_memory + output_memory + intermediate_memory + sram_memory + scaling_overhead } /// Check if FP8 precision is being used pub fn is_fp8(&self) -> bool { matches!(self.precision, PrecisionMode::FP8E4M3 { .. } | PrecisionMode::FP8E5M2 { .. }) } /// Get the FP8 config if FP8 precision is being used pub fn get_fp8_config(&self) -> Option<&FP8Config> { match &self.precision { PrecisionMode::FP8E4M3 { config } | PrecisionMode::FP8E5M2 { config } => Some(config), _ => None, } } /// Check if configuration supports training pub fn supports_training(&self) -> bool { // Training requires backward pass support and adequate precision match &self.precision { PrecisionMode::FP32 => true, PrecisionMode::FP16 => true, PrecisionMode::BF16 => true, PrecisionMode::Mixed { .. } => true, // FP8 training requires gradient precision of at least FP16 PrecisionMode::FP8E4M3 { config } | PrecisionMode::FP8E5M2 { config } => { matches!(config.gradient_precision, Precision::FP32 | Precision::FP16 | Precision::BF16) } } } /// Get unique identifier for this configuration pub fn get_identifier(&self) -> String { format!( "flash_{}heads_{}dim_{}q_{}kv_{:?}_{:?}", self.num_heads, self.head_dim, self.block_size_q, self.block_size_kv, self.memory_optimization, self.precision ) } } impl BackendConfig { /// Validate backend configuration pub fn validate(&self) -> FlashResult<()> { self.cuda.validate()?; #[cfg(feature = "edge")] if let Some(ref edge) = self.edge { edge.validate()?; } Ok(()) } } impl CudaConfig { /// Validate CUDA configuration pub fn validate(&self) -> FlashResult<()> { if self.stream_pool_size == 0 { return Err(FlashError::config("stream_pool_size must be greater than 0")); } self.memory_pool.validate()?; Ok(()) } } impl MemoryPoolConfig { /// Validate memory pool configuration pub fn validate(&self) -> FlashResult<()> { if self.initial_size == 0 { return Err(FlashError::config("memory pool initial_size must be greater than 0")); } if self.max_size < self.initial_size { return Err(FlashError::config("memory pool max_size must be >= initial_size")); } if self.growth_factor <= 1.0 { return Err(FlashError::config("memory pool growth_factor must be > 1.0")); } Ok(()) } } #[cfg(feature = "edge")] impl EdgeConfig { /// Validate edge configuration pub fn validate(&self) -> FlashResult<()> { if self.memory_limit_mb == 0 { return Err(FlashError::config("edge memory_limit_mb must be greater than 0")); } if let Some(power_limit) = self.power_limit_mw { if power_limit <= 0.0 { return Err(FlashError::config("edge power_limit_mw must be positive")); } } if let Some(latency) = self.max_latency_ms { if latency <= 0.0 { return Err(FlashError::config("edge max_latency_ms must be positive")); } } Ok(()) } } impl Default for CudaConfig { fn default() -> Self { Self { stream_pool_size: 4, enable_cuda_graphs: false, auto_tune: true, memory_pool: MemoryPoolConfig::default(), compute_capability: (8, 6), // RTX 30xx/40xx series } } } impl Default for MemoryPoolConfig { fn default() -> Self { Self { initial_size: 512 * 1024 * 1024, // 512 MB max_size: 4 * 1024 * 1024 * 1024, // 4 GB growth_factor: 1.5, enable_defragmentation: false, } } } #[cfg(test)] mod tests { use super::*; #[test] fn test_config_creation() { let config = FlashAttentionConfig::new(32, 128); assert_eq!(config.num_heads, 32); assert_eq!(config.head_dim, 128); assert_eq!(config.block_size_q, 64); assert_eq!(config.block_size_kv, 64); assert!(!config.causal); } #[test] fn test_config_validation() { let config = FlashAttentionConfig::new(32, 128); assert!(config.validate().is_ok()); let mut bad_config = config.clone(); bad_config.num_heads = 0; assert!(bad_config.validate().is_err()); bad_config = config.clone(); bad_config.head_dim = 7; // Not multiple of 8 assert!(bad_config.validate().is_err()); } #[test] fn test_softmax_scale() { let config = FlashAttentionConfig::new(32, 64); assert!((config.get_softmax_scale() - 0.125).abs() < 1e-6); } #[test] fn test_memory_estimation() { let config = FlashAttentionConfig::new(32, 128); let memory = config.estimate_memory_usage(2, 1024); assert!(memory > 0); } #[test] fn test_training_config() { let config = FlashAttentionConfig::for_training(32, 128); assert!(matches!(config.memory_optimization, MemoryOptimization::Speed)); assert!(config.backend_config.cuda.enable_cuda_graphs); } #[test] fn test_inference_config() { let config = FlashAttentionConfig::for_inference(32, 128); assert!(matches!(config.memory_optimization, MemoryOptimization::Memory)); assert!(config.backend_config.cuda.memory_pool.enable_defragmentation); } #[test] fn test_fp8_config_default() { let fp8_config = FP8Config::default(); assert!(fp8_config.use_dynamic_scaling); assert!(fp8_config.scale.is_none()); assert_eq!(fp8_config.compute_precision, Precision::FP32); assert_eq!(fp8_config.gradient_precision, Precision::BF16); assert!(fp8_config.validate().is_ok()); } #[test] fn test_fp8_inference_config() { let config = FlashAttentionConfig::for_fp8_inference(32, 128); assert!(config.is_fp8()); assert!(matches!(config.precision, PrecisionMode::FP8E4M3 { .. })); // Check compute capability requirement let (major, minor) = config.backend_config.cuda.compute_capability; assert!(major >= 8 && minor >= 9, "FP8 requires compute capability 8.9+"); // Should pass validation with correct compute capability assert!(config.validate().is_ok()); } #[test] fn test_fp8_training_config() { let config = FlashAttentionConfig::for_fp8_training(32, 128); assert!(config.is_fp8()); assert!(config.supports_training()); // Get the FP8 config let fp8_config = config.get_fp8_config().unwrap(); assert_eq!(fp8_config.compute_precision, Precision::FP32); assert_eq!(fp8_config.gradient_precision, Precision::BF16); } #[test] fn test_fp8_memory_estimation() { let fp16_config = FlashAttentionConfig::for_inference(32, 128); let fp8_config = FlashAttentionConfig::for_fp8_inference(32, 128); let fp16_memory = fp16_config.estimate_memory_usage(2, 1024); let fp8_memory = fp8_config.estimate_memory_usage(2, 1024); // FP8 should use significantly less memory for storage // Note: Compute memory may be similar due to FP16/FP32 compute assert!(fp8_memory < fp16_memory, "FP8 should use less memory than FP16"); } #[test] fn test_fp8_static_scale() { let fp8_config = FP8Config::with_static_scale(1.0); assert!(!fp8_config.use_dynamic_scaling); assert_eq!(fp8_config.scale, Some(1.0)); assert!(fp8_config.validate().is_ok()); // Invalid scale should fail validation let bad_config = FP8Config::with_static_scale(-1.0); assert!(bad_config.validate().is_err()); } #[test] fn test_precision_size_bytes() { assert_eq!(Precision::FP32.size_bytes(), 4); assert_eq!(Precision::FP16.size_bytes(), 2); assert_eq!(Precision::BF16.size_bytes(), 2); assert_eq!(Precision::FP8E4M3.size_bytes(), 1); assert_eq!(Precision::FP8E5M2.size_bytes(), 1); assert_eq!(Precision::INT8.size_bytes(), 1); } #[test] fn test_precision_is_fp8() { assert!(Precision::FP8E4M3.is_fp8()); assert!(Precision::FP8E5M2.is_fp8()); assert!(!Precision::FP16.is_fp8()); assert!(!Precision::FP32.is_fp8()); } #[test] fn test_fp8_compute_capability_requirements() { let (major, minor) = Precision::FP8E4M3.min_compute_capability(); assert_eq!(major, 8); assert_eq!(minor, 9); } #[test] fn test_fp8_validation_compute_capability() { let mut config = FlashAttentionConfig::for_fp8_inference(32, 128); // Should fail with old compute capability config.backend_config.cuda.compute_capability = (7, 5); assert!(config.validate().is_err()); // Should pass with correct compute capability config.backend_config.cuda.compute_capability = (8, 9); assert!(config.validate().is_ok()); // Should pass with higher compute capability (Hopper) config.backend_config.cuda.compute_capability = (9, 0); assert!(config.validate().is_ok()); } }