Files
rustytorch/crates/specialized/rtx-neuro-artifacts/src/explain.rs
T
2026-03-04 00:08:42 +00:00

788 lines
24 KiB
Rust

//! 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<Vec<f64>>,
/// 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<Vec<f64>>, 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<Vec<f64>> {
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<Vec<f64>> {
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<Vec<f64>> = 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<f64> = 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<f64> = (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<Vec<Vec<f64>>>,
/// 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<Vec<Vec<f64>>>, layer: usize) -> Self {
let n_heads = weights.len();
Self {
weights,
layer,
n_heads,
}
}
/// Get averaged attention across all heads
pub fn averaged(&self) -> Vec<Vec<f64>> {
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<Vec<f64>>> {
self.weights.get(head_idx)
}
}
/// Channel importance scores
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChannelImportance {
/// Importance score per channel
pub scores: Vec<f64>,
/// Channel names (if available)
pub channel_names: Option<Vec<String>>,
/// 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<f64> = saliency
.values
.iter()
.map(|ch| ch.iter().map(|v| v.abs()).sum::<f64>() / 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<String>) -> 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<f64> {
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<ArtifactType, SaliencyMap>,
/// Channel importance per artifact type
pub channel_importance: HashMap<ArtifactType, ChannelImportance>,
/// Attention maps (if available)
pub attention_maps: Option<Vec<AttentionMap>>,
/// 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<f64>],
) -> ArtifactResult<ExplanationResult> {
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<f64>],
target_type: ArtifactType,
) -> ArtifactResult<SaliencyMap> {
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<Vec<f64>> = 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<f64>],
target_type: ArtifactType,
) -> ArtifactResult<SaliencyMap> {
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<Vec<f64>> = 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<Vec<f64>> = 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<f64>]) -> Vec<Vec<f64>> {
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::<f64>()
/ (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, 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<f64> {
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<Vec<f64>> = (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);
}
}