pub mod blind_source_separation;
pub mod emd_decomposition;
pub mod multivariate_emd;
pub mod spectral_decomposition;
pub mod wavelet_transform;
use sklears_core::error::Result;
use sklears_core::types::Float;
pub use emd_decomposition::{
BoundaryCondition, EMDConfig, EMDResult, EmpiricalModeDecomposition, InterpolationMethod,
};
pub use multivariate_emd::{MEMDResult, MultivariateEMD};
pub use spectral_decomposition::{
CrossSpectralResult, SpectralDecomposition, SpectralResult, WindowFunction,
};
pub use blind_source_separation::{BSSResult, FastICA, InfoMax, NonLinearityType, JADE};
pub use wavelet_transform::{WaveletBoundary, WaveletResult, WaveletTransform, WaveletType};
pub struct SignalProcessingFactory {
default_emd_config: EMDConfig,
}
impl SignalProcessingFactory {
pub fn new() -> Self {
Self {
default_emd_config: EMDConfig::default(),
}
}
pub fn with_emd_config(config: EMDConfig) -> Self {
Self {
default_emd_config: config,
}
}
pub fn emd(&self) -> Result<EmpiricalModeDecomposition> {
let mut emd = EmpiricalModeDecomposition::new()
.tolerance(self.default_emd_config.tolerance)?
.max_sift_iter(self.default_emd_config.max_sift_iter)?
.boundary_condition(self.default_emd_config.boundary_condition)
.interpolation(self.default_emd_config.interpolation);
if let Some(max_imfs) = self.default_emd_config.max_imfs {
emd = emd.max_imfs(max_imfs)?;
}
Ok(emd)
}
pub fn emd_with_config(&self, config: EMDConfig) -> Result<EmpiricalModeDecomposition> {
let mut emd = EmpiricalModeDecomposition::new()
.tolerance(config.tolerance)?
.max_sift_iter(config.max_sift_iter)?
.boundary_condition(config.boundary_condition)
.interpolation(config.interpolation);
if let Some(max_imfs) = config.max_imfs {
emd = emd.max_imfs(max_imfs)?;
}
Ok(emd)
}
pub fn memd(&self, n_channels: usize) -> Result<MultivariateEMD> {
Ok(MultivariateEMD::new(n_channels)?.config(self.default_emd_config.clone()))
}
pub fn spectral(&self, window_size: usize) -> SpectralDecomposition {
SpectralDecomposition::new(window_size)
}
pub fn fast_ica(&self) -> FastICA {
FastICA::new()
}
pub fn jade(&self) -> JADE {
JADE::new()
}
pub fn infomax(&self) -> InfoMax {
InfoMax::new()
}
pub fn wavelet(&self, wavelet_type: WaveletType, levels: usize) -> WaveletTransform {
WaveletTransform::new()
.wavelet_type(wavelet_type)
.levels(levels)
}
}
impl Default for SignalProcessingFactory {
fn default() -> Self {
Self::new()
}
}
pub mod convenience {
use super::*;
use scirs2_core::ndarray::{Array1, Array2};
pub fn quick_emd(signal: &Array1<Float>) -> Result<EMDResult> {
let emd = EmpiricalModeDecomposition::new();
emd.decompose(signal)
}
pub fn quick_fastica(mixed_signals: &Array2<Float>) -> Result<BSSResult> {
let fastica = FastICA::new();
fastica.fit_transform(mixed_signals)
}
pub fn quick_stft(signal: &Array1<Float>, window_size: usize) -> Result<SpectralResult> {
let spectral = SpectralDecomposition::new(window_size);
spectral.stft(signal)
}
pub fn quick_wavelet(
signal: &Array1<Float>,
levels: usize,
) -> Result<crate::signal_processing::wavelet_transform::WaveletDecomposition> {
let wavelet = WaveletTransform::new()
.wavelet_type(WaveletType::Haar)
.levels(levels);
wavelet.dwt(signal).map_err(|e| {
sklears_core::error::SklearsError::Other(format!("Wavelet error: {:?}", e))
})
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::array;
#[test]
fn test_signal_processing_factory() {
let factory = SignalProcessingFactory::new();
let _emd = factory.emd().expect("valid parameter");
let _memd = factory.memd(2).expect("valid parameter");
let _spectral = factory.spectral(64);
let _fastica = factory.fast_ica();
let _jade = factory.jade();
let _infomax = factory.infomax();
let _wavelet = factory.wavelet(WaveletType::Haar, 3);
}
#[test]
fn test_convenience_functions() {
let signal = array![1.0, 2.0, 3.0, 4.0, 3.0, 2.0, 1.0, 0.0];
let _emd_result = convenience::quick_emd(&signal).expect("operation should succeed");
let _wavelet_result =
convenience::quick_wavelet(&signal, 2).expect("operation should succeed");
let _spectral_result =
convenience::quick_stft(&signal, 4).expect("operation should succeed");
}
#[test]
fn test_module_integration() {
let factory = SignalProcessingFactory::new();
let signal = array![1.0, 2.0, 3.0, 4.0, 5.0, 4.0, 3.0, 2.0, 1.0, 0.0];
let emd = factory.emd().expect("valid parameter");
let _emd_result = emd.decompose(&signal).expect("operation should succeed");
let wavelet = factory.wavelet(WaveletType::Haar, 2);
let _wavelet_result = wavelet.dwt(&signal).expect("operation should succeed");
let spectral = factory.spectral(4);
let _spectral_result = spectral.stft(&signal).expect("operation should succeed");
}
}