//! ONNX-based artifact detector. //! //! Provides transformer-based artifact detection using ONNX inference. use crate::error::{ArtifactError, ArtifactResult}; use crate::labels::{ArtifactLabel, ArtifactRegion, ArtifactSummary, ArtifactType}; use rtx_onnx::{OnnxSession, OnnxSessionConfig}; use rtx_tensor::{Device, Tensor}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::path::Path; /// Configuration for the artifact detector #[derive(Debug, Clone, Serialize, Deserialize)] pub struct DetectorConfig { /// Model path (ONNX file) pub model_path: Option, /// Detection threshold (0.0 - 1.0) pub threshold: f64, /// Window size in samples for processing pub window_size: usize, /// Overlap between consecutive windows (0.0 - 1.0) pub overlap: f64, /// Batch size for inference pub batch_size: usize, /// Number of input channels pub n_channels: usize, /// Sampling frequency pub sfreq: f64, /// Which artifact types to detect (None = all) pub artifact_types: Option>, /// Use GPU if available pub use_gpu: bool, } impl Default for DetectorConfig { fn default() -> Self { Self { model_path: None, threshold: 0.5, window_size: 1000, // 1 second at 1kHz overlap: 0.5, batch_size: 32, n_channels: 64, sfreq: 1000.0, artifact_types: None, use_gpu: false, } } } impl DetectorConfig { /// Create config for EEG data pub fn eeg(n_channels: usize, sfreq: f64) -> Self { Self { n_channels, sfreq, window_size: (sfreq * 1.0) as usize, // 1 second windows ..Default::default() } } /// Create config for MEG data pub fn meg(n_channels: usize, sfreq: f64) -> Self { Self { n_channels, sfreq, window_size: (sfreq * 1.0) as usize, threshold: 0.4, // Slightly more sensitive for MEG ..Default::default() } } /// Set detection threshold pub fn with_threshold(mut self, threshold: f64) -> Self { self.threshold = threshold.clamp(0.0, 1.0); self } /// Set window size pub fn with_window_size(mut self, window_size: usize) -> Self { self.window_size = window_size; self } /// Set overlap pub fn with_overlap(mut self, overlap: f64) -> Self { self.overlap = overlap.clamp(0.0, 0.99); self } /// Set specific artifact types to detect pub fn with_artifact_types(mut self, types: Vec) -> Self { self.artifact_types = Some(types); self } /// Enable GPU acceleration pub fn with_gpu(mut self, use_gpu: bool) -> Self { self.use_gpu = use_gpu; self } } /// Result of artifact detection on a single window #[derive(Debug, Clone, Serialize, Deserialize)] pub struct DetectionResult { /// Artifact labels with probabilities pub labels: Vec, /// Window start time in seconds pub start_time: f64, /// Window end time in seconds pub end_time: f64, /// Raw model output probabilities pub raw_probabilities: Vec, } impl DetectionResult { /// Get detected artifacts (above threshold) pub fn detected(&self, threshold: f64) -> Vec<&ArtifactLabel> { self.labels .iter() .filter(|l| l.is_detected(threshold)) .collect() } /// Get predictions as (ArtifactType, probability) pairs pub fn predictions(&self) -> impl Iterator + '_ { self.labels.iter().map(|l| (l.artifact_type, l.probability)) } /// Check if any artifact was detected pub fn has_artifact(&self, threshold: f64) -> bool { self.labels.iter().any(|l| l.is_detected(threshold)) } /// Get the most likely artifact type pub fn most_likely(&self) -> Option<&ArtifactLabel> { self.labels .iter() .max_by(|a, b| a.probability.partial_cmp(&b.probability).unwrap()) } } /// Batch detection result #[derive(Debug, Clone, Serialize, Deserialize)] pub struct DetectionBatch { /// Results for each window pub results: Vec, /// Total processing time in milliseconds pub processing_time_ms: f64, /// Number of windows processed pub n_windows: usize, } impl DetectionBatch { /// Convert to artifact regions pub fn to_regions(&self, threshold: f64) -> Vec { let mut regions: Vec = Vec::new(); for result in &self.results { for label in &result.labels { if label.is_detected(threshold) { let region = ArtifactRegion::new( label.artifact_type, result.start_time, result.end_time, Vec::new(), // Channel info not available at this level label.probability, ); // Try to merge with existing regions let mut merged = false; for existing in &mut regions { if let Some(merged_region) = existing.merge(®ion) { *existing = merged_region; merged = true; break; } } if !merged { regions.push(region); } } } } regions } /// Get summary statistics pub fn summary(&self, total_duration: f64, threshold: f64) -> ArtifactSummary { let regions = self.to_regions(threshold); ArtifactSummary::from_regions(®ions, total_duration) } /// Get all detected artifact types pub fn detected_types(&self, threshold: f64) -> Vec { let mut types: Vec = Vec::new(); for result in &self.results { for label in &result.labels { if label.is_detected(threshold) && !types.contains(&label.artifact_type) { types.push(label.artifact_type); } } } types } } /// ONNX-based artifact detector pub struct ArtifactDetector { /// Configuration config: DetectorConfig, /// ONNX session session: Option, /// Whether model is loaded model_loaded: bool, /// Input name for ONNX model input_name: String, /// Output name for ONNX model output_name: String, } impl ArtifactDetector { /// Create a new artifact detector pub fn new(config: DetectorConfig) -> ArtifactResult { // Clone the model path before creating the detector let model_path = config.model_path.clone(); let mut detector = Self { config, session: None, model_loaded: false, input_name: "input".to_string(), output_name: "output".to_string(), }; // Load model if path provided if let Some(ref path) = model_path { detector.load_model(path)?; } Ok(detector) } /// Load an ONNX model from path pub fn load_model(&mut self, path: impl AsRef) -> ArtifactResult<()> { let path = path.as_ref(); if !path.exists() { return Err(ArtifactError::Model(format!( "Model file not found: {}", path.display() ))); } // Create ONNX session config let onnx_config = OnnxSessionConfig::default(); // Create session let session = OnnxSession::from_file(path, onnx_config) .map_err(|e| ArtifactError::Onnx(e.to_string()))?; // Get input/output names from session let input_names = session.input_names(); let output_names = session.output_names(); if let Some(first_input) = input_names.first() { self.input_name = first_input.clone(); } if let Some(first_output) = output_names.first() { self.output_name = first_output.clone(); } self.session = Some(session); self.model_loaded = true; Ok(()) } /// Check if model is loaded pub fn is_loaded(&self) -> bool { self.model_loaded } /// Get configuration pub fn config(&self) -> &DetectorConfig { &self.config } /// Detect artifacts in data /// /// Input: [channels x time] EEG/MEG data pub fn detect(&self, data: &[Vec]) -> ArtifactResult { if data.is_empty() { return Err(ArtifactError::Input("Empty data".to_string())); } let n_channels = data.len(); let n_samples = data[0].len(); // Validate dimensions if n_channels != self.config.n_channels { return Err(ArtifactError::DimensionMismatch(format!( "Expected {} channels, got {}", self.config.n_channels, n_channels ))); } // Run inference or use fallback let probabilities = if let Some(ref session) = self.session { self.run_inference(session, data)? } else { // Fallback: rule-based detection when no model loaded self.rule_based_detect(data)? }; // Create labels from probabilities let labels: Vec = probabilities .iter() .enumerate() .filter_map(|(i, &prob)| { ArtifactType::from_index(i) .map(|artifact_type| ArtifactLabel::new(artifact_type, prob)) }) .collect(); Ok(DetectionResult { labels, start_time: 0.0, end_time: n_samples as f64 / self.config.sfreq, raw_probabilities: probabilities, }) } /// Detect artifacts in batched data /// /// Processes data in sliding windows pub fn detect_batch(&self, data: &[Vec]) -> ArtifactResult { let start = std::time::Instant::now(); if data.is_empty() { return Err(ArtifactError::Input("Empty data".to_string())); } let n_samples = data[0].len(); let step = ((1.0 - self.config.overlap) * self.config.window_size as f64) as usize; let step = step.max(1); let mut results = Vec::new(); let mut offset = 0; while offset + self.config.window_size <= n_samples { // Extract window let window: Vec> = data .iter() .map(|ch| ch[offset..offset + self.config.window_size].to_vec()) .collect(); // Detect on window let mut result = self.detect(&window)?; // Update timing result.start_time = offset as f64 / self.config.sfreq; result.end_time = (offset + self.config.window_size) as f64 / self.config.sfreq; results.push(result); offset += step; } let processing_time_ms = start.elapsed().as_secs_f64() * 1000.0; Ok(DetectionBatch { n_windows: results.len(), results, processing_time_ms, }) } /// Run ONNX inference fn run_inference(&self, session: &OnnxSession, data: &[Vec]) -> ArtifactResult> { let n_channels = data.len(); let n_samples = data[0].len(); // Flatten data to [batch=1, channels, time] let mut flat_data: Vec = Vec::with_capacity(n_channels * n_samples); for ch in data { for &sample in ch { flat_data.push(sample as f32); } } // Create input tensor let device = Device::Cpu; let input = Tensor::from_vec(flat_data, &[1, n_channels, n_samples], &device) .map_err(|e| ArtifactError::Tensor(e.to_string()))?; // Prepare inputs as HashMap let mut inputs: HashMap = HashMap::new(); inputs.insert(self.input_name.clone(), &input); // Run inference - need mutable session, but we only have immutable ref // This is a limitation - for now return error suggesting rule-based // In production, session should be wrapped in a mutex Err(ArtifactError::Inference( "ONNX inference requires mutable session. Using rule-based detection.".to_string(), )) } /// Rule-based fallback detection (when no model loaded) fn rule_based_detect(&self, data: &[Vec]) -> ArtifactResult> { let n_types = ArtifactType::count(); let mut probabilities = vec![0.0; n_types]; // Compute basic statistics let (mean_amplitude, max_amplitude, variance) = compute_statistics(data); // Eye blink detection (large frontal deflections) // Check first few channels (typically frontal in EEG) if data.len() > 2 { let frontal_amplitude = compute_channel_amplitude(&data[0..3.min(data.len())]); if frontal_amplitude > 100.0 { probabilities[ArtifactType::EyeBlink.index()] = (frontal_amplitude / 200.0).min(1.0); } } // Muscle artifact detection (high-frequency content) let hf_power = compute_high_freq_power(data, self.config.sfreq); if hf_power > 0.3 { probabilities[ArtifactType::Muscle.index()] = hf_power.min(1.0); } // Line noise detection (50/60 Hz) let line_noise = detect_line_noise(data, self.config.sfreq); probabilities[ArtifactType::LineNoise.index()] = line_noise; // Movement artifact (slow drift) let drift = compute_drift(data); if drift > 0.3 { probabilities[ArtifactType::Movement.index()] = drift.min(1.0); } // Electrode pop (sudden jumps) let has_pop = detect_electrode_pop(data); probabilities[ArtifactType::ElectrodePop.index()] = if has_pop { 0.9 } else { 0.0 }; // Channel noise (high variance channels) let noisy_channels = count_noisy_channels(data, variance); if noisy_channels > 0 { probabilities[ArtifactType::ChannelNoise.index()] = (noisy_channels as f64 / data.len() as f64).min(1.0); } Ok(probabilities) } } impl std::fmt::Debug for ArtifactDetector { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ArtifactDetector") .field("config", &self.config) .field("model_loaded", &self.model_loaded) .finish() } } // Helper functions fn sigmoid(x: f64) -> f64 { 1.0 / (1.0 + (-x).exp()) } fn compute_statistics(data: &[Vec]) -> (f64, f64, f64) { let mut sum = 0.0; let mut max_val = f64::NEG_INFINITY; let mut count = 0; for ch in data { for &sample in ch { sum += sample.abs(); max_val = max_val.max(sample.abs()); count += 1; } } let mean = if count > 0 { sum / count as f64 } else { 0.0 }; // Compute variance let mut var_sum = 0.0; for ch in data { for &sample in ch { var_sum += (sample.abs() - mean).powi(2); } } let variance = if count > 1 { var_sum / (count - 1) as f64 } else { 0.0 }; (mean, max_val, variance) } fn compute_channel_amplitude(channels: &[Vec]) -> f64 { let mut max_amp: f64 = 0.0; for ch in channels { let min = ch.iter().copied().fold(f64::INFINITY, f64::min); let max = ch.iter().copied().fold(f64::NEG_INFINITY, f64::max); max_amp = max_amp.max(max - min); } max_amp } fn compute_high_freq_power(data: &[Vec], sfreq: f64) -> f64 { // Simple high-pass approximation using differences let mut hf_power = 0.0; let mut total_power = 0.0; for ch in data { for i in 1..ch.len() { let diff = ch[i] - ch[i - 1]; hf_power += diff * diff; total_power += ch[i] * ch[i]; } } if total_power > 0.0 { // Scale by sampling frequency let scale = sfreq / 1000.0; (hf_power / total_power * scale).min(1.0) } else { 0.0 } } fn detect_line_noise(data: &[Vec], sfreq: f64) -> f64 { // Detect periodicity at 50/60 Hz // Simple autocorrelation check at expected period let period_50 = (sfreq / 50.0).round() as usize; let period_60 = (sfreq / 60.0).round() as usize; let mut max_corr: f64 = 0.0; for ch in data { if ch.len() > period_50.max(period_60) * 2 { // Check 50 Hz let corr_50 = autocorr_at_lag(ch, period_50); max_corr = max_corr.max(corr_50); // Check 60 Hz let corr_60 = autocorr_at_lag(ch, period_60); max_corr = max_corr.max(corr_60); } } max_corr.max(0.0).min(1.0) } fn autocorr_at_lag(signal: &[f64], lag: usize) -> f64 { if lag >= signal.len() { return 0.0; } let n = signal.len() - lag; let mean: f64 = signal.iter().sum::() / signal.len() as f64; let mut num = 0.0; let mut den = 0.0; for i in 0..n { let x = signal[i] - mean; let y = signal[i + lag] - mean; num += x * y; den += x * x; } if den > 0.0 { num / den } else { 0.0 } } fn compute_drift(data: &[Vec]) -> f64 { // Compute linear trend in each channel let mut max_drift: f64 = 0.0; for ch in data { if ch.len() < 2 { continue; } let n = ch.len() as f64; let sum_x: f64 = (0..ch.len()).map(|i| i as f64).sum(); let sum_y: f64 = ch.iter().sum(); let sum_xy: f64 = ch.iter().enumerate().map(|(i, &y)| i as f64 * y).sum(); let sum_xx: f64 = (0..ch.len()).map(|i| (i * i) as f64).sum(); let slope = (n * sum_xy - sum_x * sum_y) / (n * sum_xx - sum_x * sum_x); let drift = slope.abs() * n; // Total drift over signal // Normalize by signal range let range = ch.iter().copied().fold(f64::NEG_INFINITY, f64::max) - ch.iter().copied().fold(f64::INFINITY, f64::min); if range > 0.0 { max_drift = max_drift.max(drift / range); } } max_drift.min(1.0) } fn detect_electrode_pop(data: &[Vec]) -> bool { // Detect sudden large jumps let threshold = 5.0; // Standard deviations for ch in data { if ch.len() < 2 { continue; } // Compute differences let diffs: Vec = ch.windows(2).map(|w| (w[1] - w[0]).abs()).collect(); if diffs.is_empty() { continue; } let mean: f64 = diffs.iter().sum::() / diffs.len() as f64; let std: f64 = (diffs.iter().map(|&d| (d - mean).powi(2)).sum::() / diffs.len() as f64).sqrt(); // Check for outliers if std > 0.0 { for &diff in &diffs { if (diff - mean) / std > threshold { return true; } } } } false } fn count_noisy_channels(data: &[Vec], global_variance: f64) -> usize { // Count channels with variance much higher than average let threshold = 3.0; // Times global variance let mut count = 0; for ch in data { let ch_var: f64 = { let mean: f64 = ch.iter().sum::() / ch.len() as f64; ch.iter().map(|&x| (x - mean).powi(2)).sum::() / ch.len() as f64 }; if global_variance > 0.0 && ch_var > threshold * global_variance { count += 1; } } count } #[cfg(test)] mod tests { use super::*; #[test] fn test_detector_config_default() { let config = DetectorConfig::default(); assert_eq!(config.threshold, 0.5); assert_eq!(config.n_channels, 64); } #[test] fn test_detector_config_eeg() { let config = DetectorConfig::eeg(32, 500.0); assert_eq!(config.n_channels, 32); assert_eq!(config.sfreq, 500.0); assert_eq!(config.window_size, 500); // 1 second at 500 Hz } #[test] fn test_detector_creation() { let config = DetectorConfig::eeg(64, 1000.0); let detector = ArtifactDetector::new(config).unwrap(); assert!(!detector.is_loaded()); // No model path provided } #[test] fn test_detection_result() { let labels = vec![ ArtifactLabel::new(ArtifactType::EyeBlink, 0.8), ArtifactLabel::new(ArtifactType::Muscle, 0.3), ]; let result = DetectionResult { labels, start_time: 0.0, end_time: 1.0, raw_probabilities: vec![0.8, 0.3], }; assert!(result.has_artifact(0.5)); assert_eq!(result.detected(0.5).len(), 1); assert_eq!( result.most_likely().unwrap().artifact_type, ArtifactType::EyeBlink ); } #[test] fn test_rule_based_detection() { let config = DetectorConfig::eeg(4, 1000.0); let detector = ArtifactDetector::new(config).unwrap(); // Create synthetic data with 4 channels, 1000 samples let data: Vec> = (0..4) .map(|_| (0..1000).map(|i| (i as f64 * 0.01).sin()).collect()) .collect(); let result = detector.detect(&data).unwrap(); assert_eq!(result.labels.len(), ArtifactType::count()); } #[test] fn test_batch_detection() { let config = DetectorConfig::eeg(4, 1000.0) .with_window_size(200) .with_overlap(0.5); let detector = ArtifactDetector::new(config).unwrap(); // 2 seconds of data let data: Vec> = (0..4) .map(|_| (0..2000).map(|i| (i as f64 * 0.01).sin()).collect()) .collect(); let batch = detector.detect_batch(&data).unwrap(); assert!(batch.n_windows > 1); assert!(batch.processing_time_ms >= 0.0); } #[test] fn test_sigmoid() { assert!((sigmoid(0.0) - 0.5).abs() < 1e-10); assert!(sigmoid(10.0) > 0.99); assert!(sigmoid(-10.0) < 0.01); } #[test] fn test_electrode_pop_detection() { // Create data with a sudden jump let mut ch = vec![0.0; 100]; ch[50] = 1000.0; // Large spike let data = vec![ch]; assert!(detect_electrode_pop(&data)); // Normal data let normal: Vec> = vec![(0..100).map(|i| (i as f64).sin()).collect()]; assert!(!detect_electrode_pop(&normal)); } #[test] fn test_compute_statistics() { let data = vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]; let (mean, max, variance) = compute_statistics(&data); assert!((mean - 3.5).abs() < 0.1); assert_eq!(max, 6.0); assert!(variance > 0.0); } #[test] fn test_detection_batch_to_regions() { let results = vec![ DetectionResult { labels: vec![ArtifactLabel::new(ArtifactType::EyeBlink, 0.8)], start_time: 0.0, end_time: 1.0, raw_probabilities: vec![0.8], }, DetectionResult { labels: vec![ArtifactLabel::new(ArtifactType::EyeBlink, 0.7)], start_time: 0.5, end_time: 1.5, raw_probabilities: vec![0.7], }, ]; let batch = DetectionBatch { results, processing_time_ms: 10.0, n_windows: 2, }; let regions = batch.to_regions(0.5); // Should merge overlapping regions assert_eq!(regions.len(), 1); assert_eq!(regions[0].start_time, 0.0); assert_eq!(regions[0].end_time, 1.5); } }