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

412 lines
13 KiB
Rust

//! Minimum Norm Estimate (MNE) and related inverse methods.
//!
//! This module implements:
//! - **MNE**: Basic minimum-norm estimate
//! - **dSPM**: Dynamic Statistical Parametric Mapping (noise-normalized)
//! - **sLORETA**: Standardized Low Resolution Electromagnetic Tomography
//!
//! ## Mathematical Background
//!
//! The forward model is: M = G * J + n
//! where M is measurements, G is gain matrix, J is source activity, n is noise.
//!
//! The MNE solution is: J = W * M
//! where W = R * G^T * (G * R * G^T + λ * C)^(-1)
//!
//! - R is source covariance (with depth weighting)
//! - C is noise covariance
//! - λ is regularization parameter
use crate::{Covariance, InverseError, InverseResult, SourceEstimate};
use nalgebra::{DMatrix, DVector};
use rtx_neuro_forward::GainMatrix;
/// Inverse solution method
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum InverseMethod {
/// Basic minimum norm estimate
Mne,
/// Dynamic SPM (noise-normalized)
Dspm,
/// Standardized LORETA (resolution-normalized)
Sloreta,
}
impl Default for InverseMethod {
fn default() -> Self {
Self::Dspm
}
}
/// MNE inverse operator
#[derive(Debug, Clone)]
pub struct MneInverse {
/// Inverse kernel [n_sources x n_channels]
kernel: DMatrix<f64>,
/// Method used
method: InverseMethod,
/// Number of sources
n_sources: usize,
/// Number of channels
n_channels: usize,
/// Whether sources have free orientation
free_orientation: bool,
/// Regularization parameter (lambda^2)
lambda2: f64,
/// Depth weighting exponent
depth: f64,
/// Loose orientation constraint
loose: f64,
/// Noise normalization vector (for dSPM/sLORETA)
noise_norm: Option<DVector<f64>>,
/// Source names/indices
source_indices: Vec<usize>,
}
impl MneInverse {
/// Create an MNE inverse operator from forward model and noise covariance
///
/// # Arguments
/// * `gain` - Forward model gain matrix [n_channels x n_sources]
/// * `noise_cov` - Noise covariance matrix
/// * `method` - Inverse method (MNE, dSPM, sLORETA)
/// * `loose` - Loose orientation constraint (0 = fixed, 1 = free)
/// * `depth` - Depth weighting exponent (0 = none, typically 0.8)
/// * `lambda2` - Regularization parameter (typically 1/SNR^2)
pub fn make_inverse(
gain: &GainMatrix,
noise_cov: &Covariance,
method: InverseMethod,
loose: f64,
depth: f64,
lambda2: f64,
) -> InverseResult<Self> {
let n_channels = gain.n_sensors();
let n_source_cols = gain.n_source_columns();
let free_orientation = gain.is_free_orientation();
if noise_cov.n_channels() != n_channels {
return Err(InverseError::DimensionMismatch(format!(
"Noise covariance has {} channels, gain has {}",
noise_cov.n_channels(),
n_channels
)));
}
// Convert gain to nalgebra matrix
let g = Self::gain_to_matrix(gain);
// Compute whitener from noise covariance
let whitener = noise_cov.compute_whitener()?;
// Whiten the gain matrix: G_w = W * G
let g_white = &whitener * &g;
// Compute source covariance (with depth weighting)
let source_cov = Self::compute_source_covariance(gain, depth, loose)?;
// Compute MNE kernel: W_mne = R * G^T * (G * R * G^T + λ * I)^(-1)
// In whitened space: W_mne = R * G_w^T * (G_w * R * G_w^T + λ * I)^(-1) * W
let grgt = &g_white * &source_cov * g_white.transpose();
let grgt_reg = grgt + DMatrix::identity(n_channels, n_channels) * lambda2;
// Invert using SVD for stability
let grgt_inv = Self::pseudo_inverse(&grgt_reg)?;
// MNE kernel (before normalization)
let kernel_raw = &source_cov * g_white.transpose() * &grgt_inv * &whitener;
// Compute normalization based on method
let (kernel, noise_norm) = match method {
InverseMethod::Mne => (kernel_raw, None),
InverseMethod::Dspm => {
// dSPM: normalize by noise standard deviation
let noise_norm = Self::compute_dspm_normalization(&kernel_raw, noise_cov)?;
let kernel = Self::apply_row_normalization(&kernel_raw, &noise_norm);
(kernel, Some(noise_norm))
}
InverseMethod::Sloreta => {
// sLORETA: normalize by resolution matrix diagonal
let noise_norm = Self::compute_sloreta_normalization(&kernel_raw, &g)?;
let kernel = Self::apply_row_normalization(&kernel_raw, &noise_norm);
(kernel, Some(noise_norm))
}
};
let n_sources = if free_orientation {
n_source_cols / 3
} else {
n_source_cols
};
Ok(Self {
kernel,
method,
n_sources,
n_channels,
free_orientation,
lambda2,
depth,
loose,
noise_norm,
source_indices: (0..n_sources).collect(),
})
}
/// Apply the inverse operator to sensor data
///
/// # Arguments
/// * `data` - Sensor data [n_channels x n_times]
///
/// # Returns
/// Source estimates [n_sources x n_times] (or [3*n_sources x n_times] for free)
pub fn apply(&self, data: &[Vec<f64>]) -> InverseResult<SourceEstimate> {
if data.len() != self.n_channels {
return Err(InverseError::DimensionMismatch(format!(
"Expected {} channels, got {}",
self.n_channels,
data.len()
)));
}
let n_times = data[0].len();
// Convert to matrix
let data_mat = DMatrix::from_fn(self.n_channels, n_times, |i, j| data[i][j]);
// Apply inverse kernel: S = K * M
let source_mat = &self.kernel * &data_mat;
// Convert back to Vec<Vec<f64>>
let n_rows = source_mat.nrows();
let source_data: Vec<Vec<f64>> = (0..n_rows)
.map(|i| (0..n_times).map(|j| source_mat[(i, j)]).collect())
.collect();
// Compute times (placeholder - should come from data)
let times: Vec<f64> = (0..n_times).map(|i| i as f64).collect();
Ok(SourceEstimate::new(
source_data,
times,
self.source_indices.clone(),
self.free_orientation,
))
}
/// Get the inverse kernel matrix
pub fn kernel(&self) -> &DMatrix<f64> {
&self.kernel
}
/// Get the inverse method
pub fn method(&self) -> InverseMethod {
self.method
}
/// Get number of sources
pub fn n_sources(&self) -> usize {
self.n_sources
}
/// Get regularization parameter
pub fn lambda2(&self) -> f64 {
self.lambda2
}
// ========== Private helper methods ==========
/// Convert GainMatrix to nalgebra DMatrix
fn gain_to_matrix(gain: &GainMatrix) -> DMatrix<f64> {
let n_rows = gain.n_sensors();
let n_cols = gain.n_source_columns();
let data = gain.data();
DMatrix::from_fn(n_rows, n_cols, |i, j| data[i][j])
}
/// Compute source covariance matrix with depth weighting
fn compute_source_covariance(
gain: &GainMatrix,
depth: f64,
_loose: f64,
) -> InverseResult<DMatrix<f64>> {
let n = gain.n_source_columns();
if depth == 0.0 {
// No depth weighting: identity
return Ok(DMatrix::identity(n, n));
}
// Compute column norms of gain matrix for depth weighting
let norms = gain.source_norms();
// Depth weighting: w_i = ||g_i||^(-depth)
// Source covariance: R = diag(w_1^2, w_2^2, ...)
let weights: Vec<f64> = if gain.is_free_orientation() {
// For free orientation, apply same weight to all 3 orientations
let n_sources = n / 3;
let mut w = Vec::with_capacity(n);
for i in 0..n_sources {
let weight = if norms[i] > 1e-15 {
norms[i].powf(-depth)
} else {
1.0
};
// Same weight for x, y, z
w.push(weight * weight);
w.push(weight * weight);
w.push(weight * weight);
}
w
} else {
norms
.iter()
.map(|&norm| {
let w = if norm > 1e-15 { norm.powf(-depth) } else { 1.0 };
w * w
})
.collect()
};
Ok(DMatrix::from_diagonal(&DVector::from_vec(weights)))
}
/// Compute pseudo-inverse using SVD
fn pseudo_inverse(matrix: &DMatrix<f64>) -> InverseResult<DMatrix<f64>> {
let svd = matrix.clone().svd(true, true);
let u = svd
.u
.ok_or_else(|| InverseError::ComputationError("SVD failed".to_string()))?;
let vt = svd
.v_t
.ok_or_else(|| InverseError::ComputationError("SVD failed".to_string()))?;
let s = svd.singular_values;
let tol = 1e-10 * s[0];
let n = s.len();
let s_inv = DMatrix::from_diagonal(&DVector::from_fn(n, |i, _| {
if s[i] > tol { 1.0 / s[i] } else { 0.0 }
}));
Ok(vt.transpose() * &s_inv * u.transpose())
}
/// Compute dSPM normalization (noise standard deviation)
fn compute_dspm_normalization(
kernel: &DMatrix<f64>,
noise_cov: &Covariance,
) -> InverseResult<DVector<f64>> {
let n_sources = kernel.nrows();
let noise_cov_mat = noise_cov.data();
// Noise variance: diag(K * C * K^T)
let kckt = kernel * noise_cov_mat * kernel.transpose();
// Extract diagonal and take sqrt
let noise_norm = DVector::from_fn(n_sources, |i, _| {
let var = kckt[(i, i)];
if var > 1e-30 { var.sqrt() } else { 1.0 }
});
Ok(noise_norm)
}
/// Compute sLORETA normalization (resolution matrix diagonal)
fn compute_sloreta_normalization(
kernel: &DMatrix<f64>,
gain: &DMatrix<f64>,
) -> InverseResult<DVector<f64>> {
let n_sources = kernel.nrows();
// Resolution matrix: R = K * G
let resolution = kernel * gain;
// Take diagonal and sqrt
let noise_norm = DVector::from_fn(n_sources, |i, _| {
let val = resolution[(i, i)];
if val > 1e-30 { val.sqrt() } else { 1.0 }
});
Ok(noise_norm)
}
/// Apply row-wise normalization to kernel
fn apply_row_normalization(kernel: &DMatrix<f64>, norm: &DVector<f64>) -> DMatrix<f64> {
let n_rows = kernel.nrows();
let n_cols = kernel.ncols();
DMatrix::from_fn(n_rows, n_cols, |i, j| kernel[(i, j)] / norm[i])
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::covariance::CovarianceType;
fn create_simple_gain() -> GainMatrix {
// Simple 4x3 gain matrix (4 sensors, 3 fixed sources)
let data = vec![
vec![1.0, 0.5, 0.2],
vec![0.5, 1.0, 0.5],
vec![0.2, 0.5, 1.0],
vec![0.1, 0.3, 0.6],
];
let names = vec![
"S1".to_string(),
"S2".to_string(),
"S3".to_string(),
"S4".to_string(),
];
GainMatrix::new(data, false, names).unwrap()
}
#[test]
fn test_make_inverse_mne() {
let gain = create_simple_gain();
let noise_cov = Covariance::identity(4, CovarianceType::Noise);
let inv =
MneInverse::make_inverse(&gain, &noise_cov, InverseMethod::Mne, 0.0, 0.0, 0.1).unwrap();
assert_eq!(inv.n_sources(), 3);
assert_eq!(inv.method(), InverseMethod::Mne);
}
#[test]
fn test_make_inverse_dspm() {
let gain = create_simple_gain();
let noise_cov = Covariance::identity(4, CovarianceType::Noise);
let inv = MneInverse::make_inverse(&gain, &noise_cov, InverseMethod::Dspm, 0.0, 0.8, 0.1)
.unwrap();
assert_eq!(inv.method(), InverseMethod::Dspm);
assert!(inv.noise_norm.is_some());
}
#[test]
fn test_apply_inverse() {
let gain = create_simple_gain();
let noise_cov = Covariance::identity(4, CovarianceType::Noise);
let inv =
MneInverse::make_inverse(&gain, &noise_cov, InverseMethod::Mne, 0.0, 0.0, 0.1).unwrap();
// Apply to simple data
let data = vec![
vec![1.0, 0.0],
vec![0.0, 1.0],
vec![0.5, 0.5],
vec![0.2, 0.8],
];
let stc = inv.apply(&data).unwrap();
assert_eq!(stc.n_sources(), 3);
assert_eq!(stc.n_times(), 2);
}
}