//! 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> { 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, low_freq: f64, high_freq: f64, sfreq: f64, order: usize, ) -> PyResult>> { 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::(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, freq: f64, sfreq: f64, order: usize, ) -> PyResult>> { 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::(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, freq: f64, sfreq: f64, order: usize, ) -> PyResult>> { 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::(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, freq: f64, sfreq: f64, bandwidth: f64, harmonics: usize, ) -> PyResult>> { 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::(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, sfreq: f64, freqs: PyReadonlyArray1, n_cycles: f64, ) -> PyResult>> { let input = data.as_slice()?; let freq_vec: Vec = 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::(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::(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, sfreq: f64, n_fft: usize, n_overlap: Option, ) -> PyResult<(Bound<'py, PyArray1>, Bound<'py, PyArray1>)> { 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::(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, sfreq: f64, bandwidth: f64, n_fft: Option, ) -> PyResult<(Bound<'py, PyArray1>, Bound<'py, PyArray1>)> { 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::(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, ) -> PyResult>> { let input = data.as_slice()?; let analytic = hilbert::hilbert(input) .map_err(|e| PyErr::new::(e.to_string()))?; // Return amplitude envelope let envelope: Vec = 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, ) -> PyResult>> { let input = data.as_slice()?; let analytic = hilbert::hilbert(input) .map_err(|e| PyErr::new::(e.to_string()))?; let phase: Vec = 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, ) -> PyResult>> { let input = data.as_slice()?; let analytic = hilbert::hilbert(input) .map_err(|e| PyErr::new::(e.to_string()))?; let amplitude: Vec = analytic.iter().map(|c| c.norm()).collect(); Ok(amplitude.into_pyarray(py)) }