Initial commit
This commit is contained in:
@@ -0,0 +1,805 @@
|
||||
//! 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<String>,
|
||||
/// 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<Vec<ArtifactType>>,
|
||||
/// 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<ArtifactType>) -> 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<ArtifactLabel>,
|
||||
/// 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<f64>,
|
||||
}
|
||||
|
||||
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<Item = (ArtifactType, f64)> + '_ {
|
||||
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<DetectionResult>,
|
||||
/// 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<ArtifactRegion> {
|
||||
let mut regions: Vec<ArtifactRegion> = 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<ArtifactType> {
|
||||
let mut types: Vec<ArtifactType> = 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<OnnxSession>,
|
||||
/// 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<Self> {
|
||||
// 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<Path>) -> 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<f64>]) -> ArtifactResult<DetectionResult> {
|
||||
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<ArtifactLabel> = 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<f64>]) -> ArtifactResult<DetectionBatch> {
|
||||
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<Vec<f64>> = 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<f64>]) -> ArtifactResult<Vec<f64>> {
|
||||
let n_channels = data.len();
|
||||
let n_samples = data[0].len();
|
||||
|
||||
// Flatten data to [batch=1, channels, time]
|
||||
let mut flat_data: Vec<f32> = 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<String, &Tensor> = 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<f64>]) -> ArtifactResult<Vec<f64>> {
|
||||
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, 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>]) -> 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<f64>], 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<f64>], 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::<f64>() / 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>]) -> 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<f64>]) -> 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<f64> = ch.windows(2).map(|w| (w[1] - w[0]).abs()).collect();
|
||||
|
||||
if diffs.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let mean: f64 = diffs.iter().sum::<f64>() / diffs.len() as f64;
|
||||
let std: f64 =
|
||||
(diffs.iter().map(|&d| (d - mean).powi(2)).sum::<f64>() / 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<f64>], 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::<f64>() / ch.len() as f64;
|
||||
ch.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / 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<Vec<f64>> = (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<Vec<f64>> = (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<f64>> = 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user