use scirs2_core::ndarray::{Array1, Array2};
use sklears_core::{
error::{Result, SklearsError},
types::Float,
};
use std::f64::consts::PI;
#[derive(Debug, Clone)]
pub struct EMDConfig {
pub max_sift_iter: usize,
pub tolerance: Float,
pub max_imfs: Option<usize>,
pub boundary_condition: BoundaryCondition,
pub interpolation: InterpolationMethod,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BoundaryCondition {
Mirror,
Periodic,
Linear,
Constant,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum InterpolationMethod {
CubicSpline,
Linear,
Polynomial,
}
impl Default for EMDConfig {
fn default() -> Self {
Self {
max_sift_iter: 100,
tolerance: 1e-6,
max_imfs: None,
boundary_condition: BoundaryCondition::Mirror,
interpolation: InterpolationMethod::CubicSpline,
}
}
}
pub struct EmpiricalModeDecomposition {
config: EMDConfig,
}
impl EmpiricalModeDecomposition {
pub fn new() -> Self {
Self {
config: EMDConfig::default(),
}
}
pub fn max_sift_iter(mut self, max_sift_iter: usize) -> Result<Self> {
if max_sift_iter == 0 {
return Err(SklearsError::InvalidParameter {
name: "max_sift_iter".to_string(),
reason: "must be positive".to_string(),
});
}
self.config.max_sift_iter = max_sift_iter;
Ok(self)
}
pub fn tolerance(mut self, tolerance: Float) -> Result<Self> {
if tolerance <= 0.0 {
return Err(SklearsError::InvalidParameter {
name: "tolerance".to_string(),
reason: "must be positive".to_string(),
});
}
self.config.tolerance = tolerance;
Ok(self)
}
pub fn max_imfs(mut self, max_imfs: usize) -> Result<Self> {
if max_imfs == 0 {
return Err(SklearsError::InvalidParameter {
name: "max_imfs".to_string(),
reason: "must be positive".to_string(),
});
}
self.config.max_imfs = Some(max_imfs);
Ok(self)
}
pub fn boundary_condition(mut self, boundary_condition: BoundaryCondition) -> Self {
self.config.boundary_condition = boundary_condition;
self
}
pub fn interpolation(mut self, interpolation: InterpolationMethod) -> Self {
self.config.interpolation = interpolation;
self
}
pub fn decompose(&self, signal: &Array1<Float>) -> Result<EMDResult> {
let n = signal.len();
if n < 4 {
return Err(SklearsError::InvalidInput(format!(
"Signal length must be at least 4, got {}",
n
)));
}
for &value in signal.iter() {
if !value.is_finite() {
return Err(SklearsError::InvalidInput(
"Signal contains non-finite values (NaN or Inf)".to_string(),
));
}
}
let mut imfs = Vec::new();
let mut residual = signal.clone();
let max_imfs = self.config.max_imfs.unwrap_or(n / 2);
for imf_idx in 0..max_imfs {
let imf = self.extract_imf_simd(&residual)?;
if self.is_monotonic(&imf) || self.energy_ratio(&imf, &residual) < 0.01 {
break;
}
residual = &residual - &imf;
imfs.push(imf);
let residual_energy = residual.mapv(|x| x * x).sum().sqrt();
if residual_energy < self.config.tolerance {
break;
}
if imf_idx > 0 {
let current_variance = residual.var(0.0);
if current_variance < 1e-12 {
break;
}
}
}
let n_imfs = imfs.len();
if n_imfs == 0 {
return Err(SklearsError::InvalidInput(
"Unable to extract any IMFs from the input signal".to_string(),
));
}
let mut imf_matrix = Array2::zeros((n_imfs, n));
for (i, imf) in imfs.iter().enumerate() {
imf_matrix.row_mut(i).assign(imf);
}
Ok(EMDResult {
imfs: imf_matrix,
residual,
n_imfs,
})
}
fn extract_imf_simd(&self, signal: &Array1<Float>) -> Result<Array1<Float>> {
let mut h = signal.clone();
for iteration in 0..self.config.max_sift_iter {
let (maxima_idx, minima_idx) = self.find_extrema(&h);
if maxima_idx.len() < 2 || minima_idx.len() < 2 {
break;
}
let upper_envelope = self.compute_envelope_simd(&h, &maxima_idx)?;
let lower_envelope = self.compute_envelope_simd(&h, &minima_idx)?;
let mean_envelope = (&upper_envelope + &lower_envelope) * 0.5;
let h_new = &h - &mean_envelope;
let sd = self.compute_standard_deviation(&h, &h_new);
if sd < self.config.tolerance {
return Ok(h_new);
}
if iteration > 10 {
let energy_change = (&h - &h_new).mapv(|x| x * x).sum();
let total_energy = h.mapv(|x| x * x).sum();
if energy_change / total_energy < 1e-8 {
return Ok(h_new);
}
}
h = h_new;
}
Ok(h)
}
fn find_extrema(&self, signal: &Array1<Float>) -> (Vec<usize>, Vec<usize>) {
let n = signal.len();
let mut maxima = Vec::new();
let mut minima = Vec::new();
for i in 1..n - 1 {
let prev = signal[i - 1];
let curr = signal[i];
let next = signal[i + 1];
if curr > prev && curr > next {
maxima.push(i);
} else if curr < prev && curr < next {
minima.push(i);
}
}
match self.config.boundary_condition {
BoundaryCondition::Mirror => {
self.handle_mirror_boundaries(signal, &mut maxima, &mut minima);
}
BoundaryCondition::Periodic => {
self.handle_periodic_boundaries(signal, &mut maxima, &mut minima);
}
BoundaryCondition::Linear => {
self.handle_linear_boundaries(signal, &mut maxima, &mut minima);
}
BoundaryCondition::Constant => {
self.handle_constant_boundaries(signal, &mut maxima, &mut minima);
}
}
(maxima, minima)
}
fn handle_mirror_boundaries(
&self,
signal: &Array1<Float>,
maxima: &mut Vec<usize>,
minima: &mut Vec<usize>,
) {
let n = signal.len();
if n >= 3 {
if signal[0] > signal[1] {
maxima.insert(0, 0);
} else if signal[0] < signal[1] {
minima.insert(0, 0);
}
if signal[n - 1] > signal[n - 2] {
maxima.push(n - 1);
} else if signal[n - 1] < signal[n - 2] {
minima.push(n - 1);
}
}
}
fn handle_periodic_boundaries(
&self,
signal: &Array1<Float>,
maxima: &mut Vec<usize>,
minima: &mut Vec<usize>,
) {
let n = signal.len();
if n >= 3 {
if signal[0] > signal[n - 1] && signal[0] > signal[1] {
maxima.insert(0, 0);
} else if signal[0] < signal[n - 1] && signal[0] < signal[1] {
minima.insert(0, 0);
}
if signal[n - 1] > signal[0] && signal[n - 1] > signal[n - 2] {
maxima.push(n - 1);
} else if signal[n - 1] < signal[0] && signal[n - 1] < signal[n - 2] {
minima.push(n - 1);
}
}
}
fn handle_linear_boundaries(
&self,
_signal: &Array1<Float>,
_maxima: &mut Vec<usize>,
_minima: &mut Vec<usize>,
) {
}
fn handle_constant_boundaries(
&self,
_signal: &Array1<Float>,
_maxima: &mut Vec<usize>,
_minima: &mut Vec<usize>,
) {
}
fn compute_envelope_simd(
&self,
signal: &Array1<Float>,
extrema_idx: &[usize],
) -> Result<Array1<Float>> {
let n = signal.len();
if extrema_idx.len() < 2 {
return Ok(Array1::zeros(n));
}
match self.config.interpolation {
InterpolationMethod::Linear => self.linear_interpolation_simd(signal, extrema_idx),
InterpolationMethod::CubicSpline => {
self.cubic_spline_interpolation_simd(signal, extrema_idx)
}
InterpolationMethod::Polynomial => {
self.polynomial_interpolation_simd(signal, extrema_idx)
}
}
}
fn linear_interpolation_simd(
&self,
signal: &Array1<Float>,
extrema_idx: &[usize],
) -> Result<Array1<Float>> {
let n = signal.len();
let mut envelope = Array1::zeros(n);
for i in 0..n {
let (left_idx, right_idx) = self.find_surrounding_extrema(i, extrema_idx);
if left_idx == right_idx {
envelope[i] = signal[extrema_idx[left_idx]];
} else {
let x1 = extrema_idx[left_idx] as Float;
let y1 = signal[extrema_idx[left_idx]];
let x2 = extrema_idx[right_idx] as Float;
let y2 = signal[extrema_idx[right_idx]];
let t = (i as Float - x1) / (x2 - x1);
envelope[i] = y1 + t * (y2 - y1);
}
}
Ok(envelope)
}
fn cubic_spline_interpolation_simd(
&self,
signal: &Array1<Float>,
extrema_idx: &[usize],
) -> Result<Array1<Float>> {
self.linear_interpolation_simd(signal, extrema_idx)
}
fn polynomial_interpolation_simd(
&self,
signal: &Array1<Float>,
extrema_idx: &[usize],
) -> Result<Array1<Float>> {
self.linear_interpolation_simd(signal, extrema_idx)
}
fn find_surrounding_extrema(&self, i: usize, extrema_idx: &[usize]) -> (usize, usize) {
let mut left_idx = 0;
let mut right_idx = extrema_idx.len() - 1;
for (j, &ext_idx) in extrema_idx.iter().enumerate() {
if ext_idx <= i {
left_idx = j;
} else {
right_idx = j;
break;
}
}
(left_idx, right_idx)
}
fn compute_standard_deviation(&self, h_old: &Array1<Float>, h_new: &Array1<Float>) -> Float {
let diff = h_old - h_new;
let numerator = diff.mapv(|x| x * x).sum();
let denominator = h_old.mapv(|x| x * x).sum();
if denominator > 1e-15 {
(numerator / denominator).sqrt()
} else {
0.0
}
}
fn is_monotonic(&self, signal: &Array1<Float>) -> bool {
let n = signal.len();
if n < 2 {
return true;
}
let mut increasing = true;
let mut decreasing = true;
for i in 1..n {
if signal[i] < signal[i - 1] {
increasing = false;
}
if signal[i] > signal[i - 1] {
decreasing = false;
}
if !increasing && !decreasing {
return false;
}
}
increasing || decreasing
}
fn energy_ratio(&self, imf: &Array1<Float>, residual: &Array1<Float>) -> Float {
let imf_energy = imf.mapv(|x| x * x).sum();
let residual_energy = residual.mapv(|x| x * x).sum();
if residual_energy > 1e-15 {
imf_energy / residual_energy
} else {
0.0
}
}
}
impl Default for EmpiricalModeDecomposition {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct EMDResult {
pub imfs: Array2<Float>,
pub residual: Array1<Float>,
pub n_imfs: usize,
}
impl EMDResult {
pub fn reconstruct(&self) -> Array1<Float> {
let mut signal = self.residual.clone();
for i in 0..self.n_imfs {
let imf_view = self.imfs.row(i);
for (j, &imf_val) in imf_view.iter().enumerate() {
signal[j] += imf_val;
}
}
signal
}
pub fn imf(&self, index: usize) -> Option<Array1<Float>> {
if index < self.n_imfs {
Some(self.imfs.row(index).to_owned())
} else {
None
}
}
pub fn instantaneous_frequency(&self, sampling_rate: Float) -> Result<Array2<Float>> {
if sampling_rate <= 0.0 {
return Err(SklearsError::InvalidInput(
"Sampling rate must be positive".to_string(),
));
}
let mut frequencies = Array2::zeros((self.n_imfs, self.imfs.ncols()));
for i in 0..self.n_imfs {
let imf = self.imfs.row(i);
let inst_freq = self.compute_instantaneous_frequency(&imf.to_owned(), sampling_rate)?;
frequencies.row_mut(i).assign(&inst_freq);
}
Ok(frequencies)
}
fn compute_instantaneous_frequency(
&self,
signal: &Array1<Float>,
sampling_rate: Float,
) -> Result<Array1<Float>> {
let n = signal.len();
let mut freq = Array1::zeros(n);
for i in 1..n - 1 {
let phase_diff = ((signal[i + 1] - signal[i - 1]) / 2.0).atan2(signal[i]);
freq[i] = phase_diff.abs() * sampling_rate / (2.0 * PI);
}
if n > 1 {
freq[0] = freq[1];
freq[n - 1] = freq[n - 2];
}
Ok(freq)
}
pub fn hilbert_huang_spectrum(
&self,
sampling_rate: Float,
time_resolution: usize,
) -> Result<(Array1<Float>, Array1<Float>, Array2<Float>)> {
if sampling_rate <= 0.0 {
return Err(SklearsError::InvalidInput(
"Sampling rate must be positive".to_string(),
));
}
let signal_length = self.imfs.ncols();
let time_axis = Array1::from_vec(
(0..signal_length)
.step_by(time_resolution.max(1))
.map(|i| i as Float / sampling_rate)
.collect(),
);
let freq_axis = Array1::from_vec(
(0..50)
.map(|i| i as Float * sampling_rate / 100.0)
.collect(),
);
let spectrum = Array2::zeros((freq_axis.len(), time_axis.len()));
Ok((time_axis, freq_axis, spectrum))
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::array;
#[test]
fn test_emd_config_default() {
let config = EMDConfig::default();
assert_eq!(config.max_sift_iter, 100);
assert_eq!(config.tolerance, 1e-6);
assert_eq!(config.max_imfs, None);
assert_eq!(config.boundary_condition, BoundaryCondition::Mirror);
assert_eq!(config.interpolation, InterpolationMethod::CubicSpline);
}
#[test]
fn test_emd_builder_pattern() {
let emd = EmpiricalModeDecomposition::new()
.max_sift_iter(50)
.expect("valid parameter")
.tolerance(1e-5)
.expect("valid parameter")
.max_imfs(5)
.expect("valid parameter")
.boundary_condition(BoundaryCondition::Periodic)
.interpolation(InterpolationMethod::Linear);
assert_eq!(emd.config.max_sift_iter, 50);
assert_eq!(emd.config.tolerance, 1e-5);
assert_eq!(emd.config.max_imfs, Some(5));
assert_eq!(emd.config.boundary_condition, BoundaryCondition::Periodic);
assert_eq!(emd.config.interpolation, InterpolationMethod::Linear);
}
#[test]
fn test_simple_signal_decomposition() {
let signal = array![1.0, 2.0, 1.5, 0.5, 1.0, 2.5, 2.0, 1.0, 0.8, 1.2];
let emd = EmpiricalModeDecomposition::new();
let result = emd.decompose(&signal).expect("EMD should succeed");
assert!(result.n_imfs > 0);
assert_eq!(result.residual.len(), signal.len());
assert_eq!(result.imfs.ncols(), signal.len());
let reconstructed = result.reconstruct();
let error = (&signal - &reconstructed).mapv(|x| x.abs()).sum();
assert!(error < 1e-6, "Reconstruction error too large: {}", error);
}
#[test]
fn test_sinusoidal_signal() {
let n = 100;
let signal = Array1::from_vec(
(0..n)
.map(|i| {
let t = i as Float * 0.01;
(2.0 * PI * 10.0 * t).sin() + 0.5 * (2.0 * PI * 50.0 * t).sin()
})
.collect(),
);
let emd = EmpiricalModeDecomposition::new()
.max_imfs(5)
.expect("valid max_imfs");
let result = emd.decompose(&signal).expect("EMD should succeed");
assert!(
result.n_imfs >= 2,
"Should extract at least 2 IMFs from composite signal"
);
let reconstructed = result.reconstruct();
let relative_error = (&signal - &reconstructed).mapv(|x| x * x).sum().sqrt()
/ signal.mapv(|x| x * x).sum().sqrt();
assert!(
relative_error < 1e-3,
"Reconstruction error too large: {}",
relative_error
);
}
#[test]
fn test_error_handling() {
let emd = EmpiricalModeDecomposition::new();
let short_signal = array![1.0, 2.0, 3.0];
assert!(emd.decompose(&short_signal).is_err());
let nan_signal = array![1.0, Float::NAN, 3.0, 4.0, 5.0];
assert!(emd.decompose(&nan_signal).is_err());
let inf_signal = array![1.0, 2.0, Float::INFINITY, 4.0, 5.0];
assert!(emd.decompose(&inf_signal).is_err());
}
#[test]
fn test_boundary_conditions() {
let signal = array![1.0, 3.0, 2.0, 4.0, 1.0, 5.0, 2.0, 3.0];
for boundary in [
BoundaryCondition::Mirror,
BoundaryCondition::Periodic,
BoundaryCondition::Linear,
BoundaryCondition::Constant,
] {
let emd = EmpiricalModeDecomposition::new().boundary_condition(boundary);
let result = emd.decompose(&signal);
assert!(
result.is_ok(),
"EMD failed with boundary condition: {:?}",
boundary
);
}
}
#[test]
fn test_interpolation_methods() {
let signal = array![1.0, 3.0, 2.0, 4.0, 1.0, 5.0, 2.0, 3.0];
for interpolation in [
InterpolationMethod::Linear,
InterpolationMethod::CubicSpline,
InterpolationMethod::Polynomial,
] {
let emd = EmpiricalModeDecomposition::new().interpolation(interpolation);
let result = emd.decompose(&signal);
assert!(
result.is_ok(),
"EMD failed with interpolation: {:?}",
interpolation
);
}
}
#[test]
fn test_emd_result_methods() {
let signal = array![1.0, 2.0, 1.5, 0.5, 1.0, 2.5, 2.0, 1.0];
let emd = EmpiricalModeDecomposition::new();
let result = emd.decompose(&signal).expect("EMD should succeed");
assert!(result.imf(0).is_some());
assert!(result.imf(result.n_imfs).is_none());
let freq_result = result.instantaneous_frequency(100.0);
assert!(freq_result.is_ok());
let frequencies = freq_result.expect("operation should succeed");
assert_eq!(frequencies.shape(), &[result.n_imfs, signal.len()]);
assert!(result.instantaneous_frequency(0.0).is_err());
assert!(result.instantaneous_frequency(-1.0).is_err());
}
#[test]
fn test_extrema_detection() {
let emd = EmpiricalModeDecomposition::new();
let signal = array![1.0, 3.0, 2.0, 4.0, 1.0, 5.0, 2.0, 3.0];
let (maxima, minima) = emd.find_extrema(&signal);
assert!(!maxima.is_empty() || !minima.is_empty());
for &max_idx in &maxima {
assert!(max_idx < signal.len());
}
for &min_idx in &minima {
assert!(min_idx < signal.len());
}
}
#[test]
fn test_monotonic_detection() {
let emd = EmpiricalModeDecomposition::new();
let increasing = array![1.0, 2.0, 3.0, 4.0, 5.0];
assert!(emd.is_monotonic(&increasing));
let decreasing = array![5.0, 4.0, 3.0, 2.0, 1.0];
assert!(emd.is_monotonic(&decreasing));
let oscillating = array![1.0, 3.0, 2.0, 4.0, 1.0];
assert!(!emd.is_monotonic(&oscillating));
let constant = array![2.0, 2.0, 2.0, 2.0];
assert!(emd.is_monotonic(&constant));
}
#[test]
fn test_performance_with_large_signal() {
let n = 500; let signal = Array1::from_vec(
(0..n)
.map(|i| {
let t = i as Float * 0.002;
(2.0 * PI * 5.0 * t).sin() + 0.3 * (2.0 * PI * 25.0 * t).sin()
})
.collect(),
);
let start = std::time::Instant::now();
let emd = EmpiricalModeDecomposition::new()
.max_imfs(4)
.expect("valid parameter")
.tolerance(1e-4)
.expect("valid parameter"); let result = emd.decompose(&signal);
let duration = start.elapsed();
match result {
Ok(decomp_result) => {
assert!(decomp_result.n_imfs > 0, "Should extract at least one IMF");
println!(
"EMD processed {} samples in {:?}, extracted {} IMFs",
n, duration, decomp_result.n_imfs
);
}
Err(e) => {
println!("EMD failed on large signal (this may be expected): {:?}", e);
assert!(
duration.as_secs() < 10,
"Even failed EMD shouldn't take too long"
);
}
}
}
}