cleanup(fflag,dead): collapse flash_attention flags to unconditional — [DEAD-004]

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
This commit is contained in:
jgrusewski
2026-04-20 23:45:21 +02:00
parent 647ff0997d
commit 298cc33fcb

View File

@@ -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<GpuTensor, MLError> {
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<GpuTensor, MLError> {
// 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(())
}