use chfft::RFft1D;
use num_complex::*;
pub struct BlockConvolver<const BUFSIZE: usize> {
ir_freq_domain: Vec<Complex<f32>>,
in_freq_domain: Vec<Complex<f32>>,
fft: RFft1D<f32>,
tmp_in: Vec<f32>,
tmp_out: Vec<f32>,
remainder: Vec<f32>,
len: usize,
}
impl<const BUFSIZE: usize> std::clone::Clone for BlockConvolver<BUFSIZE> {
fn clone(&self) -> Self {
let fft = RFft1D::<f32>::new(self.len);
BlockConvolver {
ir_freq_domain: self.ir_freq_domain.clone(),
in_freq_domain: self.in_freq_domain.clone(),
fft,
tmp_in: vec![0.0; BUFSIZE * 2],
tmp_out: vec![0.0; BUFSIZE * 2],
remainder: vec![0.0; BUFSIZE],
len: self.len,
}
}
}
impl<const BUFSIZE: usize> BlockConvolver<BUFSIZE> {
pub fn from_ir(ir: &[f32]) -> Self {
let mut fft = RFft1D::<f32>::new(BUFSIZE * 2);
let mut ir_zeropad = vec![0.0; BUFSIZE * 2];
ir_zeropad[..ir.len()].copy_from_slice(ir);
BlockConvolver {
ir_freq_domain: fft.forward(&ir_zeropad),
in_freq_domain: vec![Complex::new(0.0, 0.0); ir.len() * 2],
fft,
tmp_in: vec![0.0; BUFSIZE * 2],
tmp_out: vec![0.0; BUFSIZE * 2],
remainder: vec![0.0; BUFSIZE],
len: BUFSIZE * 2,
}
}
pub fn convolve(&mut self, input: [f32; BUFSIZE]) -> [f32; BUFSIZE] {
self.tmp_in[..BUFSIZE].copy_from_slice(&self.remainder[..BUFSIZE]);
self.tmp_in[BUFSIZE..(2 * BUFSIZE)].copy_from_slice(&input[..BUFSIZE]);
self.in_freq_domain = self.fft.forward(&self.tmp_in);
for i in 0..self.in_freq_domain.len() {
self.in_freq_domain[i] = self.ir_freq_domain[i] * self.in_freq_domain[i];
}
self.tmp_out = self.fft.backward(&self.in_freq_domain);
let mut outarr = [0.0; BUFSIZE];
self.remainder[..BUFSIZE].copy_from_slice(&input[..BUFSIZE]);
outarr[..BUFSIZE].copy_from_slice(&self.tmp_out[BUFSIZE..(2 * BUFSIZE)]);
outarr
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::f32::consts::PI;
#[test]
fn test_block_convolver_freq_domain_impulse_convolution() {
let mut ir = vec![0.0; 128];
ir[0] = 1.0;
let mut signal_in = [0.0; 128];
let mut conv = BlockConvolver::<128>::from_ir(&ir);
let mut dev_accum = 0.0;
for b in 0..100 {
for i in 0..128 {
let pi_idx = ((b * 128 + i) as f32) * PI;
signal_in[i] = ((220.0 / 44100.0) * pi_idx).sin();
signal_in[i] += ((432.0 / 44100.0) * pi_idx).sin();
signal_in[i] += ((648.0 / 44100.0) * pi_idx).sin();
}
let signal_out = conv.convolve(signal_in);
for i in 0..128 {
dev_accum += (signal_out[i] - signal_in[i]) * (signal_out[i] - signal_in[i]);
}
}
assert_approx_eq::assert_approx_eq!(dev_accum / (100.0 * 128.0), 0.0, 0.00001);
}
}