use crate::{AttentionMask, Result, SparseError};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct KernelConfig {
pub hidden_dim: usize,
pub num_heads: usize,
pub head_dim: usize,
pub seq_len: usize,
pub use_flash: bool,
pub memory_format: MemoryFormat,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum MemoryFormat {
Contiguous,
ChannelsLast,
Grouped,
}
impl Default for KernelConfig {
fn default() -> Self {
Self {
hidden_dim: 2048,
num_heads: 32,
head_dim: 64,
seq_len: 4096,
use_flash: true,
memory_format: MemoryFormat::Contiguous,
}
}
}
#[derive(Debug)]
pub struct SparseKernel {
config: KernelConfig,
index_cache: std::collections::HashMap<u64, IndexMapping>,
}
#[derive(Debug, Clone)]
struct IndexMapping {
active_heads: Vec<usize>,
#[allow(dead_code)]
scatter_indices: Vec<usize>,
#[allow(dead_code)]
pattern_hash: u64,
}
impl SparseKernel {
pub fn new(config: KernelConfig) -> Self {
Self {
config,
index_cache: std::collections::HashMap::new(),
}
}
pub fn prepare(&mut self, mask: &AttentionMask, layer: usize) -> Result<()> {
let pattern_hash = self.compute_pattern_hash(mask, layer);
self.index_cache.entry(pattern_hash).or_insert_with(|| {
let active_heads = mask.active_heads(layer);
let scatter_indices: Vec<usize> =
active_heads.iter().enumerate().map(|(i, _)| i).collect();
IndexMapping {
active_heads,
scatter_indices,
pattern_hash,
}
});
Ok(())
}
pub fn execute(
&self,
mask: &AttentionMask,
layer: usize,
_q: &[f32], _k: &[f32],
_v: &[f32],
) -> Result<Vec<f32>> {
let pattern_hash = self.compute_pattern_hash(mask, layer);
let mapping = self
.index_cache
.get(&pattern_hash)
.ok_or_else(|| SparseError::KernelError("Mask pattern not prepared".into()))?;
let batch_size = 1; let output_size =
batch_size * self.config.seq_len * self.config.num_heads * self.config.head_dim;
let mut output = vec![0.0f32; output_size];
for &head in &mapping.active_heads {
let offset = head * self.config.head_dim;
for d in 0..self.config.head_dim {
if offset + d < output.len() {
output[offset + d] = 1.0; }
}
}
Ok(output)
}
fn compute_pattern_hash(&self, mask: &AttentionMask, layer: usize) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
layer.hash(&mut hasher);
for head in 0..mask.num_heads {
mask.is_active(layer, head).hash(&mut hasher);
}
hasher.finish()
}
pub fn estimate_savings(&self, mask: &AttentionMask) -> ComputeEstimate {
let total_heads = mask.num_heads * mask.num_layers;
let active_heads: usize = (0..mask.num_layers).map(|l| mask.active_count(l)).sum();
let compute_ratio = active_heads as f32 / total_heads as f32;
let memory_ratio = compute_ratio * 0.9 + 0.1;
let attention_ratio = compute_ratio;
ComputeEstimate {
compute_ratio,
memory_ratio,
attention_ratio,
estimated_speedup: 1.0 / compute_ratio,
active_heads,
total_heads,
}
}
pub fn config(&self) -> &KernelConfig {
&self.config
}
pub fn clear_cache(&mut self) {
self.index_cache.clear();
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ComputeEstimate {
pub compute_ratio: f32,
pub memory_ratio: f32,
pub attention_ratio: f32,
pub estimated_speedup: f32,
pub active_heads: usize,
pub total_heads: usize,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct KernelStats {
pub executions: u64,
pub cache_hits: u64,
pub avg_sparsity: f32,
pub compute_saved: f64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kernel_prepare() {
let mut kernel = SparseKernel::new(KernelConfig::default());
let mask = AttentionMask::random(32, 10, 0.5);
kernel.prepare(&mask, 0).unwrap();
kernel.prepare(&mask, 5).unwrap();
assert!(!kernel.index_cache.is_empty());
}
#[test]
fn test_compute_estimate() {
let kernel = SparseKernel::new(KernelConfig::default());
let mask = AttentionMask::random(32, 10, 0.5);
let estimate = kernel.estimate_savings(&mask);
assert!(estimate.compute_ratio > 0.4 && estimate.compute_ratio < 0.7);
assert!(estimate.estimated_speedup > 1.4 && estimate.estimated_speedup < 2.5);
}
#[test]
fn test_execute() {
let mut kernel = SparseKernel::new(KernelConfig {
num_heads: 8,
head_dim: 64,
seq_len: 16,
..Default::default()
});
let mask = AttentionMask::random(8, 4, 0.5);
kernel.prepare(&mask, 0).unwrap();
let q = vec![0.0f32; 16 * 8 * 64];
let k = vec![0.0f32; 16 * 8 * 64];
let v = vec![0.0f32; 16 * 8 * 64];
let output = kernel.execute(&mask, 0, &q, &k, &v).unwrap();
assert!(!output.is_empty());
}
}