From 298cc33fcbc12f89f84c050b2742b737479409ec Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 20 Apr 2026 23:45:21 +0200 Subject: [PATCH] =?UTF-8?q?cleanup(fflag,dead):=20collapse=20flash=5Fatten?= =?UTF-8?q?tion=20flags=20to=20unconditional=20=E2=80=94=20[DEAD-004]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit FlashAttention3Config had four flags, all dead or with dead else-branches: - use_sparse_patterns: write-only (sparse_pattern Mask is created unconditionally via create_sparse_mask) - io_aware_tiling: always-true setter; the "else" branch called standard_attention which itself discarded all its QK/scale/mask work and called io_aware.compute_attention — pure dead code - cuda_optimization: load_kernels() gate, always true in practice - standard_attention method + mask parameter on forward(): entirely dead Per user directive "all features enabled" / "should be used": - Deleted 4 fields (use_sparse_patterns, io_aware_tiling, cuda_optimization, sparse_pattern_iterations) — note sparse_pattern (BlockSparsePattern) stays - Collapsed forward() to unconditional io_aware.compute_attention, dropped mask param - Removed 40-LOC standard_attention dead fallback - Dropped AttentionStats.io_aware_enabled field + test assertion - cuda_kernels load unconditionally --- crates/ml/src/flash_attention/mod.rs | 70 ++-------------------------- 1 file changed, 4 insertions(+), 66 deletions(-) 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(()) }