use std::num::NonZeroUsize;
use ndarray::Array1;
use non_empty_slice::{NonEmptySlice, NonEmptyVec};
use num_complex::Complex;
use crate::Sample;
use crate::error::{SpectrogramError, SpectrogramResult};
use crate::fft_backend::C2cPlan;
use crate::spectrogram::{fft, irfft};
pub fn fft_convolve<T: Sample>(
a: &NonEmptySlice<T>,
b: &NonEmptySlice<T>,
) -> SpectrogramResult<NonEmptyVec<T>> {
let out_len = a.len().get() + b.len().get() - 1;
let n_fft = out_len.next_power_of_two();
let n = unsafe { NonZeroUsize::new_unchecked(n_fft) };
let a_spec = fft(a, n)?;
let b_spec = fft(b, n)?;
let product: Array1<Complex<T>> = &a_spec * &b_spec;
let product_slice = product.as_slice().expect("fft output is contiguous");
let product_ne = unsafe { NonEmptySlice::new_unchecked(product_slice) };
let full = irfft(product_ne, n)?;
let mut v = full.into_vec();
v.truncate(out_len);
Ok(unsafe { NonEmptyVec::new_unchecked(v) })
}
pub fn fft_deconvolve<T: Sample>(
numerator: &NonEmptySlice<T>,
denominator: &NonEmptySlice<T>,
regularization: f64,
) -> SpectrogramResult<NonEmptyVec<T>> {
let n_len = numerator.len().get();
let d_len = denominator.len().get();
let n_fft = n_len.max(d_len).next_power_of_two();
let n = unsafe { NonZeroUsize::new_unchecked(n_fft) };
let num_spec = fft(numerator, n)?;
let den_spec = fft(denominator, n)?;
let max_d2 = den_spec
.iter()
.map(Complex::norm_sqr)
.fold(T::zero(), T::max);
let eps = T::from_f64(regularization) * max_d2;
let quotient: Array1<Complex<T>> =
Array1::from_iter(num_spec.iter().zip(den_spec.iter()).map(|(nn, dd)| {
let denom = dd.norm_sqr() + eps;
if denom == T::zero() {
Complex::new(T::zero(), T::zero())
} else {
(*nn) * dd.conj() / denom
}
}));
let q_slice = quotient
.as_slice()
.expect("Array1 from_iter is always contiguous");
let q_ne = unsafe { NonEmptySlice::new_unchecked(q_slice) };
let full = irfft(q_ne, n)?;
let out_len = if n_len >= d_len {
n_len - d_len + 1
} else {
n_len
};
let mut v = full.into_vec();
v.truncate(out_len.max(1));
Ok(unsafe { NonEmptyVec::new_unchecked(v) })
}
pub struct OverlapSaveConvolver<T: Sample = f64> {
block: usize,
n_fft: usize,
overlap: usize,
fft: T::C2cPlan,
h_spec: Vec<Complex<T>>,
history: Vec<T>,
work: Vec<Complex<T>>,
inv_n: T,
}
impl<T: Sample> OverlapSaveConvolver<T> {
pub fn new(ir: &[T], block: NonZeroUsize) -> SpectrogramResult<Self> {
if ir.is_empty() {
return Err(SpectrogramError::invalid_input(
"impulse response must not be empty",
));
}
let block = block.get();
let n_fft = (block + ir.len() - 1).next_power_of_two();
let overlap = n_fft - block;
let mut fft = T::plan_c2c(n_fft)?;
let mut h_spec = vec![Complex::new(T::zero(), T::zero()); n_fft];
for (dst, &src) in h_spec.iter_mut().zip(ir.iter()) {
dst.re = src;
}
fft.forward(&mut h_spec)?;
Ok(Self {
block,
n_fft,
overlap,
fft,
h_spec,
history: vec![T::zero(); overlap],
work: vec![Complex::new(T::zero(), T::zero()); n_fft],
inv_n: T::one() / T::from_usize(n_fft),
})
}
#[must_use]
pub const fn block_size(&self) -> usize {
self.block
}
#[must_use]
pub const fn fft_size(&self) -> usize {
self.n_fft
}
pub fn reset(&mut self) {
self.history.iter_mut().for_each(|x| *x = T::zero());
}
pub fn process_block(&mut self, input: &[T], output: &mut [T]) -> SpectrogramResult<()> {
if input.len() != self.block || output.len() != self.block {
return Err(SpectrogramError::invalid_input(format!(
"process_block expects input and output of length {} (got {} and {})",
self.block,
input.len(),
output.len()
)));
}
for (dst, &src) in self.work[..self.overlap]
.iter_mut()
.zip(self.history.iter())
{
dst.re = src;
dst.im = T::zero();
}
for (dst, &src) in self.work[self.overlap..].iter_mut().zip(input.iter()) {
dst.re = src;
dst.im = T::zero();
}
for (dst, src) in self.history.iter_mut().zip(self.work[self.block..].iter()) {
*dst = src.re;
}
self.fft.forward(&mut self.work)?;
for (w, h) in self.work.iter_mut().zip(self.h_spec.iter()) {
*w *= *h;
}
self.fft.inverse(&mut self.work)?;
let inv_n = self.inv_n;
for (out, w) in output.iter_mut().zip(self.work[self.overlap..].iter()) {
*out = w.re * inv_n;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use non_empty_slice::NonEmptyVec;
fn ne(v: Vec<f64>) -> NonEmptyVec<f64> {
NonEmptyVec::new(v).unwrap()
}
#[test]
fn convolve_with_unit_impulse_returns_input_shifted() {
let signal = ne(vec![1.0, 2.0, 3.0, 4.0]);
let impulse = ne(vec![0.0, 0.0, 1.0]);
let out = fft_convolve(signal.as_non_empty_slice(), impulse.as_non_empty_slice()).unwrap();
let out = out.into_vec();
assert_eq!(out.len(), 6);
let expected = [0.0, 0.0, 1.0, 2.0, 3.0, 4.0];
for (got, want) in out.iter().zip(expected.iter()) {
assert!((got - want).abs() < 1e-9, "got {got}, want {want}");
}
}
#[test]
fn deconvolve_recovers_impulse_response() {
let x = ne(vec![1.0, 0.7, -0.3, 0.2, 0.9, -0.5, 0.1, 0.4]);
let h = ne(vec![0.0, 0.0, 1.0, 0.5]); let y = fft_convolve(x.as_non_empty_slice(), h.as_non_empty_slice()).unwrap();
let recovered = fft_deconvolve(y.as_non_empty_slice(), x.as_non_empty_slice(), 0.0)
.unwrap()
.into_vec();
let hv = h.into_vec();
assert!(recovered.len() >= hv.len());
for (i, &want) in hv.iter().enumerate() {
assert!(
(recovered[i] - want).abs() < 1e-6,
"tap {i}: got {}, want {want}",
recovered[i]
);
}
}
#[test]
fn convolve_matches_direct_convolution() {
let a = ne(vec![1.0, -2.0, 0.5]);
let b = ne(vec![0.25, 1.0, -0.5, 2.0]);
let out = fft_convolve(a.as_non_empty_slice(), b.as_non_empty_slice())
.unwrap()
.into_vec();
let av = a.into_vec();
let bv = b.into_vec();
let mut want = vec![0.0; av.len() + bv.len() - 1];
for (i, &x) in av.iter().enumerate() {
for (j, &y) in bv.iter().enumerate() {
want[i + j] += x * y;
}
}
assert_eq!(out.len(), want.len());
for (got, w) in out.iter().zip(want.iter()) {
assert!((got - w).abs() < 1e-9, "got {got}, want {w}");
}
}
#[test]
fn overlap_save_matches_direct_streaming() {
use crate::convolution::OverlapSaveConvolver;
use std::num::NonZeroUsize;
let taps = 200;
let ir: Vec<f32> = (0..taps)
.map(|k| ((k as f32 * 0.13).sin()) * (-(k as f32) / 60.0).exp())
.collect();
let total = 1024usize;
let x: Vec<f32> = (0..total)
.map(|n| (n as f32 * 0.05).sin() + 0.3 * (n as f32 * 0.21).cos())
.collect();
let block = 128usize;
let mut conv = OverlapSaveConvolver::new(&ir, NonZeroUsize::new(block).unwrap()).unwrap();
let mut got = vec![0.0f32; total];
let mut obuf = vec![0.0f32; block];
let mut start = 0;
while start + block <= total {
conv.process_block(&x[start..start + block], &mut obuf)
.unwrap();
got[start..start + block].copy_from_slice(&obuf);
start += block;
}
for n in 0..total {
let mut acc = 0.0f32;
for (k, &h) in ir.iter().enumerate() {
if n >= k {
acc += h * x[n - k];
}
}
assert!(
(got[n] - acc).abs() < 1e-3,
"sample {n}: overlap-save {} vs direct {acc}",
got[n]
);
}
}
}