use super::{EMDConfig, EmpiricalModeDecomposition};
use scirs2_core::ndarray::{s, Array1, Array2, Axis};
use sklears_core::{
error::{Result, SklearsError},
types::Float,
};
#[derive(Debug, Clone)]
pub struct MultivariateEMD {
config: EMDConfig,
n_channels: usize,
}
impl MultivariateEMD {
pub fn new(n_channels: usize) -> Result<Self> {
if n_channels == 0 {
return Err(SklearsError::InvalidParameter {
name: "n_channels".to_string(),
reason: "must be greater than 0".to_string(),
});
}
Ok(Self {
config: EMDConfig::default(),
n_channels,
})
}
pub fn config(mut self, config: EMDConfig) -> Self {
self.config = config;
self
}
pub fn decompose(&self, signals: &Array2<Float>) -> Result<MEMDResult> {
let (n_channels, n_samples) = signals.dim();
if n_channels != self.n_channels {
return Err(SklearsError::InvalidInput(format!(
"Expected {} channels, got {}",
self.n_channels, n_channels
)));
}
if n_samples < 4 {
return Err(SklearsError::InvalidInput(
"Signal length must be at least 4 samples for EMD".to_string(),
));
}
for (ch, channel) in signals.axis_iter(Axis(0)).enumerate() {
for (sample_idx, &value) in channel.iter().enumerate() {
if !value.is_finite() {
return Err(SklearsError::InvalidInput(format!(
"Non-finite value found in channel {} at sample {}",
ch, sample_idx
)));
}
}
}
let mut channel_results = Vec::with_capacity(n_channels);
for i in 0..n_channels {
let channel_signal = signals.row(i).to_owned();
let _signal_mean = channel_signal.mean().unwrap_or(0.0);
let signal_std = channel_signal.std(0.0);
if signal_std < 1e-12 {
return Err(SklearsError::InvalidInput(format!(
"Channel {} has insufficient variation (std={:.2e})",
i, signal_std
)));
}
let mut emd = EmpiricalModeDecomposition::new()
.max_sift_iter(self.config.max_sift_iter)?
.tolerance(self.config.tolerance)?
.boundary_condition(self.config.boundary_condition)
.interpolation(self.config.interpolation);
if let Some(max_imfs) = self.config.max_imfs {
emd = emd.max_imfs(max_imfs)?;
}
match emd.decompose(&channel_signal) {
Ok(result) => channel_results.push(result),
Err(e) => {
return Err(SklearsError::InvalidInput(format!(
"EMD failed for channel {}: {}",
i, e
)));
}
}
}
let min_imfs = channel_results.iter().map(|r| r.n_imfs).min().unwrap_or(0);
if min_imfs == 0 {
return Err(SklearsError::InvalidInput(
"No IMFs extracted from any channel".to_string(),
));
}
let mut aligned_imfs = Array2::zeros((n_channels * min_imfs, n_samples));
let mut residuals = Array2::zeros((n_channels, n_samples));
for (ch, result) in channel_results.iter().enumerate() {
for imf_idx in 0..min_imfs {
let global_imf_idx = ch * min_imfs + imf_idx;
let imf = result.imfs.row(imf_idx);
aligned_imfs.row_mut(global_imf_idx).assign(&imf);
}
residuals.row_mut(ch).assign(&result.residual);
}
Ok(MEMDResult {
imfs: aligned_imfs,
residuals,
n_channels,
n_imfs_per_channel: min_imfs,
})
}
}
impl Default for MultivariateEMD {
fn default() -> Self {
Self {
config: EMDConfig::default(),
n_channels: 1,
}
}
}
#[derive(Debug, Clone)]
pub struct MEMDResult {
pub imfs: Array2<Float>,
pub residuals: Array2<Float>,
pub n_channels: usize,
pub n_imfs_per_channel: usize,
}
impl MEMDResult {
pub fn channel_imfs(&self, channel: usize) -> Result<Array2<Float>> {
if channel >= self.n_channels {
return Err(SklearsError::InvalidInput(format!(
"Channel {} out of range (max: {})",
channel,
self.n_channels - 1
)));
}
let start_idx = channel * self.n_imfs_per_channel;
let end_idx = start_idx + self.n_imfs_per_channel;
Ok(self.imfs.slice(s![start_idx..end_idx, ..]).to_owned())
}
pub fn reconstruct_channel(&self, channel: usize) -> Result<Array1<Float>> {
if channel >= self.n_channels {
return Err(SklearsError::InvalidInput(format!(
"Channel {} out of range (max: {})",
channel,
self.n_channels - 1
)));
}
let channel_imfs = self.channel_imfs(channel)?;
let mut signal = self.residuals.row(channel).to_owned();
for imf in channel_imfs.axis_iter(Axis(0)) {
signal += &imf;
}
Ok(signal)
}
pub fn cross_channel_correlation(&self, imf_index: usize) -> Result<Array2<Float>> {
if imf_index >= self.n_imfs_per_channel {
return Err(SklearsError::InvalidInput(format!(
"IMF index {} out of range (max: {})",
imf_index,
self.n_imfs_per_channel - 1
)));
}
let mut correlations = Array2::zeros((self.n_channels, self.n_channels));
for i in 0..self.n_channels {
for j in 0..self.n_channels {
let imf_i_idx = i * self.n_imfs_per_channel + imf_index;
let imf_j_idx = j * self.n_imfs_per_channel + imf_index;
let imf_i = self.imfs.row(imf_i_idx);
let imf_j = self.imfs.row(imf_j_idx);
let correlation = Self::pearson_correlation(&imf_i.to_owned(), &imf_j.to_owned());
correlations[[i, j]] = correlation;
}
}
Ok(correlations)
}
pub fn imf_energy_distribution(&self) -> Array2<Float> {
let mut energies = Array2::zeros((self.n_channels, self.n_imfs_per_channel));
for channel in 0..self.n_channels {
let channel_imfs = self
.channel_imfs(channel)
.expect("operation should succeed");
let total_energy: Float = channel_imfs
.axis_iter(Axis(0))
.map(|imf| imf.mapv(|x| x * x).sum())
.sum();
if total_energy > 0.0 {
for (imf_idx, imf) in channel_imfs.axis_iter(Axis(0)).enumerate() {
let imf_energy = imf.mapv(|x| x * x).sum();
energies[[channel, imf_idx]] = imf_energy / total_energy;
}
}
}
energies
}
fn pearson_correlation(x: &Array1<Float>, y: &Array1<Float>) -> Float {
let n = x.len();
if n != y.len() || n == 0 {
return 0.0;
}
let mean_x = x.mean().unwrap_or(0.0);
let mean_y = y.mean().unwrap_or(0.0);
let mut numerator = 0.0;
let mut sum_sq_x = 0.0;
let mut sum_sq_y = 0.0;
for i in 0..n {
let dx = x[i] - mean_x;
let dy = y[i] - mean_y;
numerator += dx * dy;
sum_sq_x += dx * dx;
sum_sq_y += dy * dy;
}
let denominator = (sum_sq_x * sum_sq_y).sqrt();
if denominator == 0.0 {
0.0
} else {
(numerator / denominator).clamp(-1.0, 1.0)
}
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use crate::{BoundaryCondition, InterpolationMethod};
use scirs2_core::ndarray::Array2;
use std::f64::consts::PI;
fn generate_test_signal(n_channels: usize, n_samples: usize) -> Array2<Float> {
let mut signals = Array2::zeros((n_channels, n_samples));
for (ch, mut row) in signals.axis_iter_mut(Axis(0)).enumerate() {
for (i, val) in row.iter_mut().enumerate() {
let t = i as Float / n_samples as Float;
*val = (2.0 * PI * 5.0 * t).sin() * (1.0 + ch as Float * 0.3)
+ (2.0 * PI * 15.0 * t).sin() * 0.5
+ (2.0 * PI * 2.0 * t).cos() * 0.8
+ 0.1 * t; }
}
signals
}
#[test]
fn test_multivariate_emd_creation() {
let memd = MultivariateEMD::new(3).expect("valid parameter");
assert_eq!(memd.n_channels, 3);
let memd_default = MultivariateEMD::default();
assert_eq!(memd_default.n_channels, 1);
}
#[test]
fn test_multivariate_emd_zero_channels() {
let result = MultivariateEMD::new(0);
assert!(result.is_err());
}
#[test]
fn test_basic_decomposition() {
let signals = generate_test_signal(2, 100);
let memd = MultivariateEMD::new(2).expect("valid parameter");
let result = memd.decompose(&signals);
assert!(result.is_ok());
let result = result.expect("operation should succeed");
assert_eq!(result.n_channels, 2);
assert!(result.n_imfs_per_channel > 0);
assert_eq!(result.imfs.dim().0, 2 * result.n_imfs_per_channel);
assert_eq!(result.imfs.dim().1, 100);
assert_eq!(result.residuals.dim(), (2, 100));
}
#[test]
fn test_channel_reconstruction() {
let mut signals = Array2::zeros((2, 50));
for (ch, mut row) in signals.axis_iter_mut(Axis(0)).enumerate() {
for (i, val) in row.iter_mut().enumerate() {
let t = i as Float / 10.0; *val = (t).sin() * (1.0 + ch as Float * 0.1) + 0.01 * t;
}
}
let memd = MultivariateEMD::new(2).expect("valid parameter");
let result = memd.decompose(&signals).expect("operation should succeed");
for channel in 0..2 {
let reconstructed = result
.reconstruct_channel(channel)
.expect("operation should succeed");
let original = signals.row(channel);
assert_eq!(reconstructed.len(), original.len());
let original_energy: Float = original.mapv(|x| x.powi(2)).sum();
let reconstructed_energy: Float = reconstructed.mapv(|x| x.powi(2)).sum();
let energy_ratio = (reconstructed_energy / original_energy - 1.0).abs();
assert!(
energy_ratio < 0.5,
"Energy preservation failed for channel {}: ratio = {}",
channel,
energy_ratio
);
}
}
#[test]
fn test_cross_channel_correlation() {
let signals = generate_test_signal(3, 80);
let memd = MultivariateEMD::new(3).expect("valid parameter");
let result = memd.decompose(&signals).expect("operation should succeed");
for imf_idx in 0..result.n_imfs_per_channel {
let correlations = result
.cross_channel_correlation(imf_idx)
.expect("operation should succeed");
assert_eq!(correlations.dim(), (3, 3));
for i in 0..3 {
assert!((correlations[[i, i]] - 1.0).abs() < 1e-10);
}
for i in 0..3 {
for j in 0..3 {
let diff = (correlations[[i, j]] - correlations[[j, i]]).abs();
assert!(diff < 1e-10, "Correlation matrix not symmetric");
}
}
}
}
#[test]
fn test_energy_distribution() {
let signals = generate_test_signal(2, 60);
let memd = MultivariateEMD::new(2).expect("valid parameter");
let result = memd.decompose(&signals).expect("operation should succeed");
let energies = result.imf_energy_distribution();
assert_eq!(energies.dim(), (2, result.n_imfs_per_channel));
for channel in 0..2 {
let total_energy: Float = energies.row(channel).sum();
assert!((total_energy - 1.0).abs() < 1e-6);
}
}
#[test]
fn test_error_handling() {
let memd = MultivariateEMD::new(2).expect("valid parameter");
let wrong_signals = Array2::zeros((3, 50));
let result = memd.decompose(&wrong_signals);
assert!(result.is_err());
let short_signals = Array2::zeros((2, 3));
let result = memd.decompose(&short_signals);
assert!(result.is_err());
let mut bad_signals = Array2::zeros((2, 50));
bad_signals[[0, 10]] = Float::NAN;
let result = memd.decompose(&bad_signals);
assert!(result.is_err());
}
#[test]
fn test_constant_signal_handling() {
let memd = MultivariateEMD::new(1).expect("valid parameter");
let constant_signal = Array2::ones((1, 50));
let result = memd.decompose(&constant_signal);
assert!(result.is_err()); }
#[test]
fn test_channel_access_bounds() {
let signals = generate_test_signal(2, 40);
let memd = MultivariateEMD::new(2).expect("valid parameter");
let result = memd.decompose(&signals).expect("operation should succeed");
assert!(result.channel_imfs(0).is_ok());
assert!(result.channel_imfs(1).is_ok());
assert!(result.reconstruct_channel(0).is_ok());
assert!(result.reconstruct_channel(1).is_ok());
assert!(result.channel_imfs(2).is_err());
assert!(result.reconstruct_channel(2).is_err());
if result.n_imfs_per_channel > 0 {
assert!(result.cross_channel_correlation(0).is_ok());
}
assert!(result
.cross_channel_correlation(result.n_imfs_per_channel)
.is_err());
}
#[test]
fn test_custom_configuration() {
let config = EMDConfig {
max_sift_iter: 20,
tolerance: 1e-4,
max_imfs: Some(3),
boundary_condition: BoundaryCondition::Periodic,
interpolation: InterpolationMethod::CubicSpline,
};
let signals = generate_test_signal(2, 100);
let memd = MultivariateEMD::new(2)
.expect("valid parameter")
.config(config);
let result = memd.decompose(&signals).expect("operation should succeed");
assert!(result.n_imfs_per_channel <= 3);
}
}