788 lines
24 KiB
Rust
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);
|
|
}
|
|
}
|