use rustfft::{num_complex::Complex64, FftPlanner};
fn fft_real(signal: &[f64]) -> Vec<Complex64> {
let n = signal.len();
if n == 0 {
return Vec::new();
}
let mut buffer: Vec<Complex64> = signal.iter().map(|&x| Complex64::new(x, 0.0)).collect();
let mut planner = FftPlanner::new();
let fft = planner.plan_fft_forward(n);
fft.process(&mut buffer);
buffer.truncate(n / 2 + 1);
buffer
}
fn periodogram(signal: &[f64]) -> Vec<(usize, f64)> {
let n = signal.len();
if n < 4 {
return Vec::new();
}
let fft_result = fft_real(signal);
let n_f64 = n as f64;
let mut result = Vec::with_capacity(n / 2);
for (k, complex) in fft_result.iter().enumerate().skip(1) {
let period = n / k;
if period < 2 {
break;
}
let power = (complex.re * complex.re + complex.im * complex.im) / n_f64;
result.push((period, power));
}
result.sort_by(|a, b| b.0.cmp(&a.0));
result
}
pub fn welch_periodogram(signal: &[f64], window_size: usize, overlap: f64) -> Vec<(usize, f64)> {
let n = signal.len();
if n < window_size || window_size < 4 {
return periodogram(signal);
}
let overlap = overlap.clamp(0.0, 0.9);
let hop = ((1.0 - overlap) * window_size as f64).ceil() as usize;
let hop = hop.max(1);
let mut accumulated_psd: std::collections::HashMap<usize, (f64, usize)> =
std::collections::HashMap::new();
let mut start = 0;
while start + window_size <= n {
let segment = &signal[start..start + window_size];
let windowed: Vec<f64> = segment
.iter()
.enumerate()
.map(|(i, &x)| {
let w = 0.5
* (1.0 - (2.0 * std::f64::consts::PI * i as f64 / window_size as f64).cos());
x * w
})
.collect();
let psd = periodogram(&windowed);
for (period, power) in psd {
let entry = accumulated_psd.entry(period).or_insert((0.0, 0));
entry.0 += power;
entry.1 += 1;
}
start += hop;
}
let mut result: Vec<(usize, f64)> = accumulated_psd
.into_iter()
.map(|(period, (sum, count))| (period, sum / count as f64))
.collect();
result.sort_by(|a, b| b.0.cmp(&a.0));
result
}
#[cfg(test)]
mod tests {
use super::*;
fn generate_sine(n: usize, period: usize) -> Vec<f64> {
(0..n)
.map(|i| (2.0 * std::f64::consts::PI * i as f64 / period as f64).sin())
.collect()
}
#[test]
fn welch_periodogram_basic() {
let signal = generate_sine(256, 12);
let psd = welch_periodogram(&signal, 64, 0.5);
assert!(!psd.is_empty());
let peak = psd.iter().max_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
assert!(peak.is_some());
let (period, _) = peak.unwrap();
assert!(
(10..=14).contains(period),
"Expected period near 12, got {}",
period
);
}
#[test]
fn welch_short_signal() {
let signal = generate_sine(32, 8);
let psd = welch_periodogram(&signal, 64, 0.5);
assert!(!psd.is_empty());
}
#[test]
fn welch_overlap_values() {
let signal = generate_sine(256, 16);
for overlap in [0.0, 0.25, 0.5, 0.75] {
let psd = welch_periodogram(&signal, 64, overlap);
assert!(!psd.is_empty(), "Failed with overlap {}", overlap);
}
}
#[test]
fn welch_finds_multiple_periods() {
let n = 512;
let signal: Vec<f64> = (0..n)
.map(|i| {
(2.0 * std::f64::consts::PI * i as f64 / 16.0).sin()
+ 0.5 * (2.0 * std::f64::consts::PI * i as f64 / 32.0).sin()
})
.collect();
let psd = welch_periodogram(&signal, 128, 0.5);
let top_periods: Vec<usize> = psd.iter().take(10).map(|(p, _)| *p).collect();
let has_16 = top_periods.iter().any(|p| (14..=18).contains(p));
let has_32 = top_periods.iter().any(|p| (28..=36).contains(p));
assert!(
has_16 || has_32,
"Should detect at least one period, got {:?}",
top_periods
);
}
}