diff --git a/crates/ml/src/flash_attention/mod.rs b/crates/ml/src/flash_attention/mod.rs index e59bd839d..b8798be3d 100644 --- a/crates/ml/src/flash_attention/mod.rs +++ b/crates/ml/src/flash_attention/mod.rs @@ -183,10 +183,7 @@ pub struct FlashAttention3Config { pub head_dim: usize, pub max_seq_len: usize, pub dropout_rate: f32, - pub use_sparse_patterns: bool, pub sparse_pattern: BlockSparsePattern, - pub io_aware_tiling: bool, - pub cuda_optimization: bool, } impl Default for FlashAttention3Config { @@ -197,10 +194,7 @@ impl Default for FlashAttention3Config { head_dim: 64, max_seq_len: 1024, dropout_rate: 0.1, - use_sparse_patterns: true, sparse_pattern: BlockSparsePattern::default(), - io_aware_tiling: true, - cuda_optimization: true, } } } @@ -224,10 +218,7 @@ impl FlashAttention3 { let io_aware = IOAwareAttention::new(64, 2048); // 64 tile size, 2GB memory budget let causal_optimizer = CausalMaskOptimizer::new(1024); let mut cuda_manager = CudaKernelManager::new(); - - if config.cuda_optimization { - cuda_manager.load_kernels()?; - } + cuda_manager.load_kernels()?; let ctx = cudarc::driver::CudaContext::new(0) .map_err(|e| MLError::DeviceError(format!("CUDA context: {e}")))?; @@ -249,68 +240,18 @@ impl FlashAttention3 { }) } - /// Compute attention using Flash Attention 3 + /// Compute attention using Flash Attention 3 (IO-aware tiling, unconditional). pub fn forward( &mut self, q: &GpuTensor, k: &GpuTensor, v: &GpuTensor, - mask: Option<&GpuTensor>, ) -> Result { let (_batch_size, _seq_len, _) = q .dims3() .map_err(|e| MLError::ModelError(format!("Invalid Q tensor dims: {}", e)))?; - // Use IO-aware attention for computation - let output = if self.config.io_aware_tiling { - self.io_aware.compute_attention(q, k, v)? - } else { - // Fallback to standard attention computation - self.standard_attention(q, k, v, mask)? - }; - - Ok(output) - } - - fn standard_attention( - &self, - q: &GpuTensor, - k: &GpuTensor, - v: &GpuTensor, - mask: Option<&GpuTensor>, - ) -> Result { - // Compute Q @ K^T - let k_t = k.transpose(1, 2, &self.stream) - .map_err(|e| MLError::ModelError(format!("K transpose failed: {}", e)))?; - let scores = q - .matmul(&k_t, &self.cublas, &self.stream) - .map_err(|e| MLError::ModelError(format!("QK computation failed: {}", e)))?; - - // Scale by sqrt(head_dim): divide by scalar via broadcast_div with a scalar tensor - let scale = (self.config.head_dim as f32).sqrt(); - let scale_tensor = GpuTensor::scalar(scale, &self.stream) - .map_err(|e| MLError::ModelError(format!("Scale tensor: {}", e)))?; - let scaled_scores = scores - .broadcast_div(&scale_tensor, &self.stream) - .map_err(|e| MLError::ModelError(format!("Score scaling failed: {}", e)))?; - - // Apply mask if provided - let masked_scores = if let Some(mask) = mask { - scaled_scores - .add(mask, &self.stream) - .map_err(|e| MLError::ModelError(format!("Mask application failed: {}", e)))? - } else { - scaled_scores - }; - - // Apply softmax via GPU sigmoid approximation (no ActivationKernels::softmax exists) - // For attention: use IO-aware path which handles this correctly. - // Fallback: just return V weighted by IO-aware attention. - let output = self.io_aware.compute_attention(q, k, v)?; - - let _ = masked_scores; // silence unused warning - - Ok(output) + self.io_aware.compute_attention(q, k, v) } /// Create sparse attention mask @@ -323,7 +264,6 @@ impl FlashAttention3 { AttentionStats { cache_size: self.attention_cache.len(), cuda_kernels_loaded: self.cuda_manager.kernels_loaded, - io_aware_enabled: self.config.io_aware_tiling, } } } @@ -333,7 +273,6 @@ impl FlashAttention3 { pub struct AttentionStats { pub cache_size: usize, pub cuda_kernels_loaded: bool, - pub io_aware_enabled: bool, } #[cfg(test)] @@ -377,7 +316,7 @@ mod tests { let v = GpuTensor::from_host(&v_data, vec![batch_size, seq_len, head_dim], &attention.stream) .map_err(|e| MLError::ModelError(e.to_string()))?; - let output = attention.forward(&q, &k, &v, None)?; + let output = attention.forward(&q, &k, &v)?; // Check output dimensions assert_eq!(output.dims(), q.dims()); @@ -405,7 +344,6 @@ mod tests { let stats = attention.get_stats(); assert_eq!(stats.cache_size, 0); - assert!(stats.io_aware_enabled); Ok(()) }