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

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))
}