339 lines
9.5 KiB
Rust
339 lines
9.5 KiB
Rust
//! Signal processing module
|
|
|
|
use pyo3::prelude::*;
|
|
use numpy::{PyArray1, PyArray2, PyReadonlyArray1, PyReadonlyArray2, IntoPyArray};
|
|
use rtx_neuro_signal::{filter, wavelet, psd, hilbert};
|
|
|
|
/// Create the signal submodule
|
|
pub fn create_module(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
|
|
let m = PyModule::new(py, "signal")?;
|
|
|
|
// Filtering functions
|
|
m.add_function(wrap_pyfunction!(filter_bandpass, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(filter_highpass, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(filter_lowpass, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(filter_notch, &m)?)?;
|
|
|
|
// Time-frequency
|
|
m.add_function(wrap_pyfunction!(morlet_wavelet, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(compute_psd, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(compute_psd_multitaper, &m)?)?;
|
|
|
|
// Phase analysis
|
|
m.add_function(wrap_pyfunction!(hilbert_transform, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(instantaneous_phase, &m)?)?;
|
|
m.add_function(wrap_pyfunction!(instantaneous_amplitude, &m)?)?;
|
|
|
|
Ok(m)
|
|
}
|
|
|
|
/// Apply bandpass filter to data
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D or 2D array of signal data
|
|
/// * `low_freq` - Low cutoff frequency (Hz)
|
|
/// * `high_freq` - High cutoff frequency (Hz)
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `order` - Filter order (default: 4)
|
|
///
|
|
/// # Returns
|
|
/// Filtered data with same shape as input
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, low_freq, high_freq, sfreq, order=4))]
|
|
fn filter_bandpass<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
low_freq: f64,
|
|
high_freq: f64,
|
|
sfreq: f64,
|
|
order: usize,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
let config = filter::FilterConfig {
|
|
low_freq: Some(low_freq),
|
|
high_freq: Some(high_freq),
|
|
sfreq,
|
|
order,
|
|
filter_type: filter::FilterType::Butterworth,
|
|
zero_phase: true,
|
|
};
|
|
|
|
let output = filter::bandpass(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok(output.into_pyarray(py))
|
|
}
|
|
|
|
/// Apply highpass filter
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `freq` - Cutoff frequency (Hz)
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `order` - Filter order (default: 4)
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, freq, sfreq, order=4))]
|
|
fn filter_highpass<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
freq: f64,
|
|
sfreq: f64,
|
|
order: usize,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
let config = filter::FilterConfig {
|
|
low_freq: Some(freq),
|
|
high_freq: None,
|
|
sfreq,
|
|
order,
|
|
filter_type: filter::FilterType::Butterworth,
|
|
zero_phase: true,
|
|
};
|
|
|
|
let output = filter::highpass(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok(output.into_pyarray(py))
|
|
}
|
|
|
|
/// Apply lowpass filter
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `freq` - Cutoff frequency (Hz)
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `order` - Filter order (default: 4)
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, freq, sfreq, order=4))]
|
|
fn filter_lowpass<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
freq: f64,
|
|
sfreq: f64,
|
|
order: usize,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
let config = filter::FilterConfig {
|
|
low_freq: None,
|
|
high_freq: Some(freq),
|
|
sfreq,
|
|
order,
|
|
filter_type: filter::FilterType::Butterworth,
|
|
zero_phase: true,
|
|
};
|
|
|
|
let output = filter::lowpass(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok(output.into_pyarray(py))
|
|
}
|
|
|
|
/// Apply notch filter (e.g., for line noise removal)
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `freq` - Center frequency to remove (Hz)
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `bandwidth` - Notch bandwidth (Hz, default: 1.0)
|
|
/// * `harmonics` - Number of harmonics to filter (default: 1)
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, freq, sfreq, bandwidth=1.0, harmonics=1))]
|
|
fn filter_notch<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
freq: f64,
|
|
sfreq: f64,
|
|
bandwidth: f64,
|
|
harmonics: usize,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
let config = filter::NotchConfig {
|
|
center_freq: freq,
|
|
sfreq,
|
|
bandwidth,
|
|
n_harmonics: harmonics,
|
|
};
|
|
|
|
let output = filter::notch(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok(output.into_pyarray(py))
|
|
}
|
|
|
|
/// Compute Morlet wavelet time-frequency decomposition
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `freqs` - Array of frequencies to analyze
|
|
/// * `n_cycles` - Number of cycles (default: 7.0)
|
|
///
|
|
/// # Returns
|
|
/// Complex 2D array [n_freqs x n_times]
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, sfreq, freqs, n_cycles=7.0))]
|
|
fn morlet_wavelet<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
sfreq: f64,
|
|
freqs: PyReadonlyArray1<f64>,
|
|
n_cycles: f64,
|
|
) -> PyResult<Bound<'py, PyArray2<f64>>> {
|
|
let input = data.as_slice()?;
|
|
let freq_vec: Vec<f64> = freqs.as_slice()?.to_vec();
|
|
|
|
let config = wavelet::MorletConfig {
|
|
sfreq,
|
|
n_cycles,
|
|
frequencies: freq_vec.clone(),
|
|
};
|
|
|
|
let tfr = wavelet::morlet_power(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
// Convert to 2D array [n_freqs x n_times]
|
|
let n_freqs = freq_vec.len();
|
|
let n_times = tfr.len() / n_freqs;
|
|
|
|
let array = ndarray::Array2::from_shape_vec((n_freqs, n_times), tfr)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok(array.into_pyarray(py))
|
|
}
|
|
|
|
/// Compute power spectral density using Welch's method
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `n_fft` - FFT length (default: 256)
|
|
/// * `n_overlap` - Overlap samples (default: n_fft/2)
|
|
///
|
|
/// # Returns
|
|
/// Tuple of (frequencies, psd)
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, sfreq, n_fft=256, n_overlap=None))]
|
|
fn compute_psd<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
sfreq: f64,
|
|
n_fft: usize,
|
|
n_overlap: Option<usize>,
|
|
) -> PyResult<(Bound<'py, PyArray1<f64>>, Bound<'py, PyArray1<f64>>)> {
|
|
let input = data.as_slice()?;
|
|
let overlap = n_overlap.unwrap_or(n_fft / 2);
|
|
|
|
let config = psd::WelchConfig {
|
|
sfreq,
|
|
n_fft,
|
|
n_overlap: overlap,
|
|
window: psd::Window::Hann,
|
|
};
|
|
|
|
let (freqs, power) = psd::welch(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok((freqs.into_pyarray(py), power.into_pyarray(py)))
|
|
}
|
|
|
|
/// Compute power spectral density using multitaper method
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
/// * `sfreq` - Sampling frequency (Hz)
|
|
/// * `bandwidth` - Frequency bandwidth (Hz, default: 4.0)
|
|
/// * `n_fft` - FFT length (default: None = data length)
|
|
///
|
|
/// # Returns
|
|
/// Tuple of (frequencies, psd)
|
|
#[pyfunction]
|
|
#[pyo3(signature = (data, sfreq, bandwidth=4.0, n_fft=None))]
|
|
fn compute_psd_multitaper<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
sfreq: f64,
|
|
bandwidth: f64,
|
|
n_fft: Option<usize>,
|
|
) -> PyResult<(Bound<'py, PyArray1<f64>>, Bound<'py, PyArray1<f64>>)> {
|
|
let input = data.as_slice()?;
|
|
let fft_len = n_fft.unwrap_or(input.len());
|
|
|
|
let config = psd::MultitaperConfig {
|
|
sfreq,
|
|
bandwidth,
|
|
n_fft: fft_len,
|
|
};
|
|
|
|
let (freqs, power) = psd::multitaper(input, &config)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
Ok((freqs.into_pyarray(py), power.into_pyarray(py)))
|
|
}
|
|
|
|
/// Compute Hilbert transform
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
///
|
|
/// # Returns
|
|
/// Analytic signal (complex)
|
|
#[pyfunction]
|
|
fn hilbert_transform<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
|
|
let analytic = hilbert::hilbert(input)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
// Return amplitude envelope
|
|
let envelope: Vec<f64> = analytic.iter().map(|c| c.norm()).collect();
|
|
|
|
Ok(envelope.into_pyarray(py))
|
|
}
|
|
|
|
/// Compute instantaneous phase using Hilbert transform
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
///
|
|
/// # Returns
|
|
/// Phase in radians [-π, π]
|
|
#[pyfunction]
|
|
fn instantaneous_phase<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
|
|
let analytic = hilbert::hilbert(input)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
let phase: Vec<f64> = analytic.iter().map(|c| c.arg()).collect();
|
|
|
|
Ok(phase.into_pyarray(py))
|
|
}
|
|
|
|
/// Compute instantaneous amplitude using Hilbert transform
|
|
///
|
|
/// # Arguments
|
|
/// * `data` - 1D array of signal data
|
|
///
|
|
/// # Returns
|
|
/// Amplitude envelope
|
|
#[pyfunction]
|
|
fn instantaneous_amplitude<'py>(
|
|
py: Python<'py>,
|
|
data: PyReadonlyArray1<f64>,
|
|
) -> PyResult<Bound<'py, PyArray1<f64>>> {
|
|
let input = data.as_slice()?;
|
|
|
|
let analytic = hilbert::hilbert(input)
|
|
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(e.to_string()))?;
|
|
|
|
let amplitude: Vec<f64> = analytic.iter().map(|c| c.norm()).collect();
|
|
|
|
Ok(amplitude.into_pyarray(py))
|
|
}
|