//! Explainability for artifact detection. //! //! Provides interpretable saliency maps and attention visualizations //! to explain which channels and timepoints contributed to artifact detection. use crate::detector::{ArtifactDetector, DetectionResult}; use crate::error::{ArtifactError, ArtifactResult}; use crate::labels::ArtifactType; use serde::{Deserialize, Serialize}; use std::collections::HashMap; /// Configuration for explainability #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ExplainerConfig { /// Number of steps for integrated gradients pub ig_steps: usize, /// Baseline type for integrated gradients pub baseline: BaselineType, /// Whether to compute channel importance pub compute_channel_importance: bool, /// Whether to compute temporal importance pub compute_temporal_importance: bool, /// Smoothing factor for saliency maps (0 = no smoothing) pub smoothing: f64, /// Number of samples for SHAP (if enabled) pub shap_samples: usize, } impl Default for ExplainerConfig { fn default() -> Self { Self { ig_steps: 50, baseline: BaselineType::Zero, compute_channel_importance: true, compute_temporal_importance: true, smoothing: 0.0, shap_samples: 100, } } } impl ExplainerConfig { /// Use fewer steps for faster computation pub fn fast() -> Self { Self { ig_steps: 20, shap_samples: 50, ..Default::default() } } /// Use more steps for higher accuracy pub fn accurate() -> Self { Self { ig_steps: 100, shap_samples: 200, ..Default::default() } } } /// Baseline type for integrated gradients #[derive(Debug, Clone, Copy, Serialize, Deserialize)] pub enum BaselineType { /// Zero baseline Zero, /// Mean of the input Mean, /// Gaussian noise Noise, /// Uniform random Random, } /// Saliency map for artifact detection explanation #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SaliencyMap { /// Saliency values [channels x time] pub values: Vec>, /// Target artifact type this explains pub artifact_type: ArtifactType, /// Method used to compute saliency pub method: String, /// Minimum saliency value pub min_value: f64, /// Maximum saliency value pub max_value: f64, } impl SaliencyMap { /// Create a new saliency map pub fn new(values: Vec>, artifact_type: ArtifactType, method: &str) -> Self { let (min_value, max_value) = compute_min_max(&values); Self { values, artifact_type, method: method.to_string(), min_value, max_value, } } /// Get normalized saliency (0-1 range) pub fn normalized(&self) -> Vec> { let range = self.max_value - self.min_value; if range <= 0.0 { return self.values.clone(); } self.values .iter() .map(|ch| ch.iter().map(|&v| (v - self.min_value) / range).collect()) .collect() } /// Get absolute saliency values pub fn absolute(&self) -> Vec> { self.values .iter() .map(|ch| ch.iter().map(|v| v.abs()).collect()) .collect() } /// Apply smoothing to the saliency map pub fn smoothed(&self, window_size: usize) -> Self { let smoothed: Vec> = self .values .iter() .map(|ch| smooth_signal(ch, window_size)) .collect(); Self::new(smoothed, self.artifact_type, &self.method) } /// Get channel with highest saliency pub fn most_important_channel(&self) -> usize { let channel_importance: Vec = self .values .iter() .map(|ch| ch.iter().map(|v| v.abs()).sum()) .collect(); channel_importance .iter() .enumerate() .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap()) .map_or(0, |(i, _)| i) } /// Get time point with highest saliency pub fn most_important_timepoint(&self) -> usize { let n_times = self.values.first().map_or(0, std::vec::Vec::len); if n_times == 0 { return 0; } let time_importance: Vec = (0..n_times) .map(|t| { self.values .iter() .map(|ch| ch.get(t).map_or(0.0, |v| v.abs())) .sum() }) .collect(); time_importance .iter() .enumerate() .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap()) .map_or(0, |(i, _)| i) } } /// Attention map from transformer layers #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AttentionMap { /// Attention weights [heads x query_len x key_len] pub weights: Vec>>, /// Layer index pub layer: usize, /// Number of attention heads pub n_heads: usize, } impl AttentionMap { /// Create a new attention map pub fn new(weights: Vec>>, layer: usize) -> Self { let n_heads = weights.len(); Self { weights, layer, n_heads, } } /// Get averaged attention across all heads pub fn averaged(&self) -> Vec> { if self.weights.is_empty() { return Vec::new(); } let n_queries = self.weights[0].len(); let n_keys = self.weights[0].first().map_or(0, std::vec::Vec::len); let mut avg = vec![vec![0.0; n_keys]; n_queries]; for head in &self.weights { for (i, row) in head.iter().enumerate() { for (j, &val) in row.iter().enumerate() { avg[i][j] += val / self.n_heads as f64; } } } avg } /// Get attention for a specific head pub fn head(&self, head_idx: usize) -> Option<&Vec>> { self.weights.get(head_idx) } } /// Channel importance scores #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ChannelImportance { /// Importance score per channel pub scores: Vec, /// Channel names (if available) pub channel_names: Option>, /// Artifact type this is computed for pub artifact_type: ArtifactType, } impl ChannelImportance { /// Create from saliency map pub fn from_saliency(saliency: &SaliencyMap) -> Self { let scores: Vec = saliency .values .iter() .map(|ch| ch.iter().map(|v| v.abs()).sum::() / ch.len() as f64) .collect(); Self { scores, channel_names: None, artifact_type: saliency.artifact_type, } } /// Set channel names pub fn with_names(mut self, names: Vec) -> Self { self.channel_names = Some(names); self } /// Get top N most important channels pub fn top_channels(&self, n: usize) -> Vec<(usize, f64)> { let mut indexed: Vec<(usize, f64)> = self .scores .iter() .enumerate() .map(|(i, &s)| (i, s)) .collect(); indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); indexed.truncate(n); indexed } /// Normalize scores to sum to 1 pub fn normalized(&self) -> Vec { let sum: f64 = self.scores.iter().sum(); if sum <= 0.0 { return self.scores.clone(); } self.scores.iter().map(|&s| s / sum).collect() } } /// Complete explanation result #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ExplanationResult { /// Saliency maps per artifact type pub saliency_maps: HashMap, /// Channel importance per artifact type pub channel_importance: HashMap, /// Attention maps (if available) pub attention_maps: Option>, /// Detection result being explained pub detection: DetectionResult, /// Computation time in milliseconds pub computation_time_ms: f64, } impl ExplanationResult { /// Get explanation for a specific artifact type pub fn for_artifact(&self, artifact_type: ArtifactType) -> Option<&SaliencyMap> { self.saliency_maps.get(&artifact_type) } /// Get channel importance for a specific artifact type pub fn channel_importance_for( &self, artifact_type: ArtifactType, ) -> Option<&ChannelImportance> { self.channel_importance.get(&artifact_type) } /// Get most detected artifact with its explanation pub fn most_detected(&self) -> Option<(&ArtifactType, &SaliencyMap)> { let most_detected = self .detection .labels .iter() .max_by(|a, b| a.probability.partial_cmp(&b.probability).unwrap())?; let saliency = self.saliency_maps.get(&most_detected.artifact_type)?; Some((&most_detected.artifact_type, saliency)) } } /// Explainer for artifact detection pub struct ArtifactExplainer { /// Configuration config: ExplainerConfig, } impl ArtifactExplainer { /// Create a new explainer pub fn new(config: ExplainerConfig) -> Self { Self { config } } /// Create with default configuration pub fn default_explainer() -> Self { Self::new(ExplainerConfig::default()) } /// Explain artifact detection pub fn explain( &self, detector: &ArtifactDetector, data: &[Vec], ) -> ArtifactResult { let start = std::time::Instant::now(); // First get detection results let detection = detector.detect(data)?; // Compute saliency maps for detected artifacts let mut saliency_maps = HashMap::new(); let mut channel_importance = HashMap::new(); for label in &detection.labels { if label.probability > 0.1 { // Only explain if some probability // Compute saliency using finite differences (gradient approximation) let saliency = self.compute_saliency(detector, data, label.artifact_type)?; // Compute channel importance from saliency let importance = ChannelImportance::from_saliency(&saliency); channel_importance.insert(label.artifact_type, importance); saliency_maps.insert(label.artifact_type, saliency); } } let computation_time_ms = start.elapsed().as_secs_f64() * 1000.0; Ok(ExplanationResult { saliency_maps, channel_importance, attention_maps: None, // Would require model internals detection, computation_time_ms, }) } /// Compute saliency map using finite differences fn compute_saliency( &self, detector: &ArtifactDetector, data: &[Vec], target_type: ArtifactType, ) -> ArtifactResult { let n_channels = data.len(); let n_times = data.first().map_or(0, std::vec::Vec::len); if n_channels == 0 || n_times == 0 { return Err(ArtifactError::Input("Empty data".to_string())); } // Get baseline output let baseline_result = detector.detect(data)?; let baseline_prob = baseline_result .labels .iter() .find(|l| l.artifact_type == target_type) .map_or(0.0, |l| l.probability); // Compute gradient approximation using finite differences let epsilon = 1e-5; let mut saliency = vec![vec![0.0; n_times]; n_channels]; // For efficiency, we sample a subset of positions let time_step = (n_times / 50).max(1); for ch in 0..n_channels { for t in (0..n_times).step_by(time_step) { let mut perturbed: Vec> = data.to_vec(); perturbed[ch][t] += epsilon; let perturbed_result = detector.detect(&perturbed)?; let perturbed_prob = perturbed_result .labels .iter() .find(|l| l.artifact_type == target_type) .map_or(0.0, |l| l.probability); let gradient = (perturbed_prob - baseline_prob) / epsilon; // Fill in the region around this sample let start = t.saturating_sub(time_step / 2); let end = (t + time_step / 2).min(n_times); for ti in start..end { saliency[ch][ti] = gradient; } } } // Apply smoothing if configured if self.config.smoothing > 0.0 { let window = (self.config.smoothing * 10.0) as usize; saliency = saliency .iter() .map(|ch| smooth_signal(ch, window.max(3))) .collect(); } Ok(SaliencyMap::new( saliency, target_type, "finite_differences", )) } /// Compute integrated gradients (requires model gradients) pub fn integrated_gradients( &self, detector: &ArtifactDetector, data: &[Vec], target_type: ArtifactType, ) -> ArtifactResult { let n_channels = data.len(); let n_times = data.first().map_or(0, std::vec::Vec::len); // Generate baseline let baseline = self.create_baseline(data); // Compute integrated gradients along path from baseline to input let mut accumulated_grads = vec![vec![0.0; n_times]; n_channels]; for step in 0..self.config.ig_steps { let alpha = step as f64 / self.config.ig_steps as f64; // Interpolate between baseline and input let interpolated: Vec> = data .iter() .zip(baseline.iter()) .map(|(d, b)| { d.iter() .zip(b.iter()) .map(|(&di, &bi)| bi + alpha * (di - bi)) .collect() }) .collect(); // Get detection at this point let result = detector.detect(&interpolated)?; let prob = result .labels .iter() .find(|l| l.artifact_type == target_type) .map_or(0.0, |l| l.probability); // Approximate gradient let epsilon = 1e-5; for ch in 0..n_channels.min(10) { // Sample channels for efficiency for t in (0..n_times).step_by(n_times / 20 + 1) { let mut perturbed = interpolated.clone(); perturbed[ch][t] += epsilon; let perturbed_result = detector.detect(&perturbed)?; let perturbed_prob = perturbed_result .labels .iter() .find(|l| l.artifact_type == target_type) .map_or(0.0, |l| l.probability); accumulated_grads[ch][t] += (perturbed_prob - prob) / epsilon; } } } // Scale by (input - baseline) and normalize by number of steps let saliency: Vec> = accumulated_grads .iter() .enumerate() .map(|(ch, grads)| { grads .iter() .enumerate() .map(|(t, &g)| { let diff = data[ch][t] - baseline[ch][t]; g * diff / self.config.ig_steps as f64 }) .collect() }) .collect(); Ok(SaliencyMap::new( saliency, target_type, "integrated_gradients", )) } /// Create baseline based on configuration fn create_baseline(&self, data: &[Vec]) -> Vec> { let n_channels = data.len(); let n_times = data.first().map_or(0, std::vec::Vec::len); match self.config.baseline { BaselineType::Zero => vec![vec![0.0; n_times]; n_channels], BaselineType::Mean => { let global_mean: f64 = data.iter().flat_map(|ch| ch.iter()).sum::() / (n_channels * n_times) as f64; vec![vec![global_mean; n_times]; n_channels] } BaselineType::Noise => { // Simple pseudo-random noise let mut baseline = vec![vec![0.0; n_times]; n_channels]; let mut seed = 12345u64; for ch in &mut baseline { for val in ch.iter_mut() { seed = seed.wrapping_mul(1103515245).wrapping_add(12345); *val = (seed as f64 / u64::MAX as f64) * 0.1 - 0.05; } } baseline } BaselineType::Random => { // Uniform random in data range let (min_val, max_val) = compute_min_max(data); let range = max_val - min_val; let mut baseline = vec![vec![0.0; n_times]; n_channels]; let mut seed = 54321u64; for ch in &mut baseline { for val in ch.iter_mut() { seed = seed.wrapping_mul(1103515245).wrapping_add(12345); *val = min_val + (seed as f64 / u64::MAX as f64) * range; } } baseline } } } } impl std::fmt::Debug for ArtifactExplainer { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ArtifactExplainer") .field("config", &self.config) .finish() } } // Helper functions fn compute_min_max(data: &[Vec]) -> (f64, f64) { let mut min_val = f64::INFINITY; let mut max_val = f64::NEG_INFINITY; for ch in data { for &val in ch { min_val = min_val.min(val); max_val = max_val.max(val); } } (min_val, max_val) } fn smooth_signal(signal: &[f64], window_size: usize) -> Vec { if signal.is_empty() || window_size <= 1 { return signal.to_vec(); } let half_window = window_size / 2; let mut smoothed = Vec::with_capacity(signal.len()); for i in 0..signal.len() { let start = i.saturating_sub(half_window); let end = (i + half_window + 1).min(signal.len()); let sum: f64 = signal[start..end].iter().sum(); smoothed.push(sum / (end - start) as f64); } smoothed } #[cfg(test)] mod tests { use super::*; use crate::detector::DetectorConfig; #[test] fn test_explainer_config_default() { let config = ExplainerConfig::default(); assert_eq!(config.ig_steps, 50); assert!(config.compute_channel_importance); } #[test] fn test_explainer_config_fast() { let config = ExplainerConfig::fast(); assert!(config.ig_steps < ExplainerConfig::default().ig_steps); } #[test] fn test_saliency_map_creation() { let values = vec![vec![0.1, 0.5, 0.3], vec![0.2, 0.8, 0.1]]; let saliency = SaliencyMap::new(values, ArtifactType::EyeBlink, "test"); assert_eq!(saliency.artifact_type, ArtifactType::EyeBlink); assert_eq!(saliency.min_value, 0.1); assert_eq!(saliency.max_value, 0.8); } #[test] fn test_saliency_normalized() { let values = vec![vec![0.0, 0.5, 1.0]]; let saliency = SaliencyMap::new(values, ArtifactType::Muscle, "test"); let norm = saliency.normalized(); assert_eq!(norm[0][0], 0.0); assert!((norm[0][2] - 1.0).abs() < 1e-10); } #[test] fn test_saliency_most_important_channel() { let values = vec![ vec![0.1, 0.1, 0.1], // Channel 0: sum = 0.3 vec![0.5, 0.5, 0.5], // Channel 1: sum = 1.5 vec![0.2, 0.2, 0.2], // Channel 2: sum = 0.6 ]; let saliency = SaliencyMap::new(values, ArtifactType::Muscle, "test"); assert_eq!(saliency.most_important_channel(), 1); } #[test] fn test_saliency_most_important_timepoint() { let values = vec![ vec![0.1, 0.9, 0.1], // Sum at t=1 is highest vec![0.1, 0.8, 0.1], ]; let saliency = SaliencyMap::new(values, ArtifactType::EyeBlink, "test"); assert_eq!(saliency.most_important_timepoint(), 1); } #[test] fn test_saliency_smoothed() { let values = vec![vec![1.0, 0.0, 1.0, 0.0, 1.0]]; let saliency = SaliencyMap::new(values, ArtifactType::Muscle, "test"); let smoothed = saliency.smoothed(3); // Smoothing should reduce variance let orig_var: f64 = saliency.values[0].iter().map(|&v| (v - 0.6).powi(2)).sum(); let smooth_var: f64 = smoothed.values[0].iter().map(|&v| (v - 0.6).powi(2)).sum(); assert!(smooth_var <= orig_var); } #[test] fn test_attention_map() { let weights = vec![ vec![vec![0.5, 0.5], vec![0.3, 0.7]], // Head 0 vec![vec![0.4, 0.6], vec![0.6, 0.4]], // Head 1 ]; let attn = AttentionMap::new(weights, 0); assert_eq!(attn.n_heads, 2); let avg = attn.averaged(); assert_eq!(avg.len(), 2); assert!((avg[0][0] - 0.45).abs() < 1e-10); } #[test] fn test_channel_importance_from_saliency() { let values = vec![ vec![0.1, 0.2, 0.3], // Channel 0: mean = 0.2 vec![0.4, 0.5, 0.6], // Channel 1: mean = 0.5 ]; let saliency = SaliencyMap::new(values, ArtifactType::Muscle, "test"); let importance = ChannelImportance::from_saliency(&saliency); assert_eq!(importance.scores.len(), 2); assert!(importance.scores[1] > importance.scores[0]); } #[test] fn test_channel_importance_top_channels() { let importance = ChannelImportance { scores: vec![0.1, 0.5, 0.3, 0.8, 0.2], channel_names: None, artifact_type: ArtifactType::EyeBlink, }; let top = importance.top_channels(3); assert_eq!(top.len(), 3); assert_eq!(top[0].0, 3); // Highest score assert_eq!(top[1].0, 1); assert_eq!(top[2].0, 2); } #[test] fn test_explainer_creation() { let config = ExplainerConfig::default(); let explainer = ArtifactExplainer::new(config); assert_eq!(explainer.config.ig_steps, 50); } #[test] fn test_explainer_explain() { let detector_config = DetectorConfig::eeg(4, 100.0); let detector = ArtifactDetector::new(detector_config).unwrap(); let data: Vec> = (0..4) .map(|_| (0..100).map(|i| (i as f64 * 0.1).sin()).collect()) .collect(); let explainer = ArtifactExplainer::new(ExplainerConfig::fast()); let result = explainer.explain(&detector, &data).unwrap(); assert!(result.computation_time_ms >= 0.0); assert!(!result.detection.labels.is_empty()); } #[test] fn test_smooth_signal() { let signal = vec![1.0, 0.0, 1.0, 0.0, 1.0]; let smoothed = smooth_signal(&signal, 3); assert_eq!(smoothed.len(), signal.len()); // Middle values should be smoothed toward average assert!(smoothed[2] < 1.0); assert!(smoothed[2] > 0.0); } #[test] fn test_baseline_types() { let data = vec![vec![1.0, 2.0, 3.0]; 2]; let config = ExplainerConfig::default(); let explainer = ArtifactExplainer::new(config); // Zero baseline let mut explainer_zero = ArtifactExplainer::new(ExplainerConfig { baseline: BaselineType::Zero, ..Default::default() }); let baseline = explainer_zero.create_baseline(&data); assert!(baseline[0].iter().all(|&v| v == 0.0)); // Mean baseline let mut explainer_mean = ArtifactExplainer::new(ExplainerConfig { baseline: BaselineType::Mean, ..Default::default() }); let baseline = explainer_mean.create_baseline(&data); assert!((baseline[0][0] - 2.0).abs() < 1e-10); } }