use crate::complex::Complex;
use crate::error::{MathError, Result};
use std::f64::consts::PI;
fn check_n(n: usize) -> Result<()> {
if n == 0 {
return Err(MathError::InvalidArgument("FFT size must be > 0".into()));
}
if n & (n - 1) != 0 {
return Err(MathError::InvalidArgument(format!(
"FFT size must be a power of 2, got {}",
n
)));
}
Ok(())
}
pub fn fft(input: &[Complex<f64>]) -> Result<Vec<Complex<f64>>> {
let n = input.len();
check_n(n)?;
let mut buf = input.to_vec();
fft_in_place(&mut buf, false);
Ok(buf)
}
pub fn ifft(input: &[Complex<f64>]) -> Result<Vec<Complex<f64>>> {
let n = input.len();
check_n(n)?;
let mut buf = input.to_vec();
fft_in_place(&mut buf, true);
for x in &mut buf {
x.re /= n as f64;
x.im /= n as f64;
}
Ok(buf)
}
pub fn rfft(input: &[f64]) -> Result<Vec<Complex<f64>>> {
let n = input.len();
check_n(n)?;
let complex_input: Vec<Complex<f64>> =
input.iter().map(|&x| Complex::new(x, 0.0)).collect();
let full = fft(&complex_input)?;
Ok(full[..n / 2 + 1].to_vec())
}
pub fn magnitude_spectrum(samples: &[f64]) -> Result<Vec<f64>> {
Ok(rfft(samples)?.iter().map(|c| c.abs()).collect())
}
pub fn power_spectrum(samples: &[f64]) -> Result<Vec<f64>> {
Ok(rfft(samples)?.iter().map(|c| c.abs().powi(2)).collect())
}
pub fn fft_in_place(buf: &mut [Complex<f64>], inverse: bool) {
let n = buf.len();
debug_assert!(n > 0 && (n & (n - 1)) == 0, "FFT size must be a power of 2");
let bits = n.trailing_zeros() as usize;
for i in 0..n {
let j = reverse_bits(i, bits);
if i < j {
buf.swap(i, j);
}
}
let sign = if inverse { 1.0 } else { -1.0 };
let mut size = 2;
while size <= n {
let half = size / 2;
let theta = sign * 2.0 * std::f64::consts::PI / size as f64;
let twiddles: Vec<Complex<f64>> = (0..half)
.map(|k| {
let angle = theta * k as f64;
Complex::new(angle.cos(), angle.sin())
})
.collect();
let mut start = 0;
while start < n {
for k in 0..half {
let t = twiddles[k] * buf[start + k + half];
let u = buf[start + k];
buf[start + k] = u + t;
buf[start + k + half] = u - t;
}
start += size;
}
size *= 2;
}
}
fn reverse_bits(x: usize, bits: usize) -> usize {
const REV_BYTE: [u8; 256] = {
let mut table = [0u8; 256];
let mut i = 0;
while i < 256 {
let mut val = i as u8;
let mut rev = 0u8;
let mut j = 0;
while j < 8 {
rev = (rev << 1) | (val & 1);
val >>= 1;
j += 1;
}
table[i] = rev;
i += 1;
}
table
};
let b0 = REV_BYTE[x & 0xFF] as usize;
let b1 = REV_BYTE[(x >> 8) & 0xFF] as usize;
let b2 = REV_BYTE[(x >> 16) & 0xFF] as usize;
let b3 = REV_BYTE[(x >> 24) & 0xFF] as usize;
let b4 = REV_BYTE[(x >> 32) & 0xFF] as usize;
let b5 = REV_BYTE[(x >> 40) & 0xFF] as usize;
let b6 = REV_BYTE[(x >> 48) & 0xFF] as usize;
let b7 = REV_BYTE[(x >> 56) & 0xFF] as usize;
let reversed = (b0 << 56) | (b1 << 48) | (b2 << 40) | (b3 << 32)
| (b4 << 24) | (b5 << 16) | (b6 << 8) | b7;
reversed >> (64 - bits)
}
pub fn fft2(input: &[Vec<Complex<f64>>]) -> Result<Vec<Vec<Complex<f64>>>> {
if input.is_empty() {
return Ok(Vec::new());
}
let rows = input.len();
let cols = input[0].len();
check_n(rows)?;
check_n(cols)?;
let mut data: Vec<Vec<Complex<f64>>> = input.to_vec();
for r in &mut data {
let mut row = std::mem::take(r);
fft_in_place(&mut row, false);
*r = row;
}
for c in 0..cols {
let mut col: Vec<Complex<f64>> = (0..rows).map(|r| data[r][c]).collect();
fft_in_place(&mut col, false);
for r in 0..rows {
data[r][c] = col[r];
}
}
Ok(data)
}
pub fn next_pow2(n: usize) -> usize {
if n <= 1 {
return 1;
}
1 << (usize::BITS - (n - 1).leading_zeros())
}
const HAMMING_A0: f64 = 0.54;
const HAMMING_A1: f64 = 0.46;
const BLACKMAN_A0: f64 = 0.42;
const BLACKMAN_A1: f64 = 0.5;
const BLACKMAN_A2: f64 = 0.08;
pub fn apply_window(samples: &[f64], window: Window) -> Vec<f64> {
let n = samples.len();
if n <= 1 {
return samples.to_vec();
}
let denom = (n - 1) as f64;
match window {
Window::Rectangular => samples.to_vec(),
Window::Hann => samples
.iter()
.enumerate()
.map(|(i, &x)| x * 0.5 * (1.0 - (2.0 * PI * i as f64 / denom).cos()))
.collect(),
Window::Hamming => samples
.iter()
.enumerate()
.map(|(i, &x)| x * (HAMMING_A0 - HAMMING_A1 * (2.0 * PI * i as f64 / denom).cos()))
.collect(),
Window::Blackman => samples
.iter()
.enumerate()
.map(|(i, &x)| {
let t = 2.0 * PI * i as f64 / denom;
x * (BLACKMAN_A0 - BLACKMAN_A1 * t.cos() + BLACKMAN_A2 * (2.0 * t).cos())
})
.collect(),
}
}
#[derive(Debug, Clone, Copy)]
pub enum Window {
Rectangular,
Hann,
Hamming,
Blackman,
}
pub fn convolve(a: &[f64], b: &[f64]) -> Result<Vec<f64>> {
if a.is_empty() || b.is_empty() {
return Err(MathError::InvalidArgument("convolve: empty input".into()));
}
let result_len = a.len() + b.len() - 1;
let n = next_pow2(result_len);
let mut fa: Vec<Complex<f64>> = a.iter().map(|&x| Complex::new(x, 0.0)).collect();
let mut fb: Vec<Complex<f64>> = b.iter().map(|&x| Complex::new(x, 0.0)).collect();
fa.resize(n, Complex::ZERO);
fb.resize(n, Complex::ZERO);
fft_in_place(&mut fa, false);
fft_in_place(&mut fb, false);
for i in 0..n {
fa[i] = fa[i] * fb[i];
}
fft_in_place(&mut fa, true);
let scale = 1.0 / n as f64;
Ok(fa[..result_len].iter().map(|c| c.re * scale).collect())
}
pub fn cross_correlate(a: &[f64], b: &[f64]) -> Result<Vec<f64>> {
if a.is_empty() || b.is_empty() {
return Err(MathError::InvalidArgument("cross_correlate: empty input".into()));
}
let b_rev: Vec<f64> = b.iter().rev().copied().collect();
convolve(a, &b_rev)
}
#[cfg(test)]
mod tests {
use super::*;
use std::f64::consts::PI;
fn close(a: f64, b: f64, eps: f64) -> bool {
(a - b).abs() < eps
}
#[test]
fn dft_impulse() {
let x = vec![
Complex::new(1.0, 0.0),
Complex::ZERO,
Complex::ZERO,
Complex::ZERO,
];
let y = fft(&x).unwrap();
for v in &y {
assert!(close(v.re, 1.0, 1e-12));
assert!(close(v.im, 0.0, 1e-12));
}
}
#[test]
fn dft_complex_exponential() {
let n = 8;
let x: Vec<Complex<f64>> = (0..n)
.map(|i| Complex::from_polar(1.0, 2.0 * PI * i as f64 / n as f64))
.collect();
let y = fft(&x).unwrap();
for (k, v) in y.iter().enumerate() {
if k == 1 {
assert!(close(v.abs(), n as f64, 1e-9));
} else {
assert!(close(v.abs(), 0.0, 1e-9));
}
}
}
#[test]
fn inverse_roundtrip() {
let x: Vec<Complex<f64>> = (0..16)
.map(|i| Complex::new(i as f64, -(i as f64) * 0.5))
.collect();
let y = fft(&x).unwrap();
let back = ifft(&y).unwrap();
for (a, b) in x.iter().zip(back.iter()) {
assert!(close(a.re, b.re, 1e-9));
assert!(close(a.im, b.im, 1e-9));
}
}
#[test]
fn real_sine_magnitude() {
let n = 64;
let samples: Vec<f64> = (0..n).map(|i| (2.0 * PI * 4.0 * i as f64 / n as f64).sin()).collect();
let mags = magnitude_spectrum(&samples).unwrap();
assert!(mags[4] > 30.0, "expected strong bin at k=4, got {}", mags[4]);
for (k, m) in mags.iter().enumerate() {
if k != 4 && k != (n - 4) && k != 0 {
assert!(*m < 1.0, "unexpected energy at k={}: {}", k, m);
}
}
}
#[test]
fn power_of_two_check() {
assert!(fft(&[Complex::ZERO; 3]).is_err());
assert!(fft(&[Complex::ZERO; 8]).is_ok());
}
#[test]
fn next_pow2_basic() {
assert_eq!(next_pow2(1), 1);
assert_eq!(next_pow2(5), 8);
assert_eq!(next_pow2(8), 8);
assert_eq!(next_pow2(9), 16);
}
#[test]
fn fft2_roundtrip() {
let n = 4;
let mut data = vec![vec![Complex::ZERO; n]; n];
for r in 0..n {
for c in 0..n {
data[r][c] = Complex::new((r * n + c) as f64, 0.0);
}
}
let transformed = fft2(&data).unwrap();
assert_eq!(transformed.len(), n);
assert_eq!(transformed[0].len(), n);
}
#[test]
fn convolve_impulse() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![1.0];
let c = convolve(&a, &b).unwrap();
assert_eq!(c.len(), 4);
for i in 0..4 {
assert!(close(c[i], a[i], 1e-9));
}
}
#[test]
fn convolve_known() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 1.0];
let c = convolve(&a, &b).unwrap();
assert_eq!(c.len(), 4);
assert!(close(c[0], 1.0, 1e-9));
assert!(close(c[1], 3.0, 1e-9));
assert!(close(c[2], 5.0, 1e-9));
assert!(close(c[3], 3.0, 1e-9));
}
#[test]
fn cross_correlate_self() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let c = cross_correlate(&a, &a).unwrap();
assert_eq!(c.len(), 7);
let peak_idx = c
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
.map(|(i, _)| i)
.unwrap();
assert_eq!(peak_idx, 3);
assert!(close(c[3], 30.0, 1e-8));
}
#[test]
fn window_hann_endpoints() {
let n = 8;
let samples = vec![1.0; n];
let w = apply_window(&samples, Window::Hann);
assert!(close(w[0], 0.0, 1e-10));
assert!(close(w[n - 1], 0.0, 1e-10));
assert!(w[n / 2] > 0.9 && w[n / 2] < 1.01);
}
#[test]
fn window_hamming_endpoints() {
let n = 8;
let samples = vec![1.0; n];
let w = apply_window(&samples, Window::Hamming);
assert!(close(w[0], 0.08, 1e-10));
assert!(close(w[n - 1], 0.08, 1e-10));
assert!(w[n / 2] > 0.9 && w[n / 2] < 1.01);
}
#[test]
fn window_blackman_endpoints() {
let n = 8;
let samples = vec![1.0; n];
let w = apply_window(&samples, Window::Blackman);
assert!(close(w[0], 0.0, 1e-10));
assert!(close(w[n - 1], 0.0, 1e-10));
}
#[test]
fn window_rectangular_identity() {
let samples = vec![1.0, 2.0, 3.0, 4.0];
let w = apply_window(&samples, Window::Rectangular);
assert_eq!(w, samples);
}
}