use nalgebra::DVector;
use num_complex::Complex;
use rustfft::{Fft, FftPlanner};
use std::sync::Arc;
pub struct FdafAec {
fft_size: usize,
frame_size: usize,
fft: Arc<dyn Fft<f32>>,
ifft: Arc<dyn Fft<f32>>,
weights: DVector<Complex<f32>>,
far_end_buffer: DVector<f32>,
mu: f32,
psd: DVector<f32>,
smoothing_factor: f32,
}
impl FdafAec {
pub fn new(fft_size: usize, step_size: f32) -> Self {
assert!(fft_size > 0 && fft_size.is_power_of_two(), "fft_size must be a power of two.");
let frame_size = fft_size / 2;
let mut fft_planner = FftPlanner::new();
let fft = fft_planner.plan_fft_forward(fft_size);
let ifft = fft_planner.plan_fft_inverse(fft_size);
Self {
fft_size,
frame_size,
fft,
ifft,
weights: DVector::from_element(fft_size, Complex::new(0.0, 0.0)),
far_end_buffer: DVector::from_element(fft_size, 0.0),
mu: step_size,
psd: DVector::from_element(fft_size, 1.0), smoothing_factor: 0.98,
}
}
pub fn process(&mut self, far_end_frame: &[f32], mic_frame: &[f32]) -> Vec<f32> {
assert_eq!(far_end_frame.len(), self.frame_size, "Input far-end frame size must be half of FFT size.");
assert_eq!(mic_frame.len(), self.frame_size, "Input mic frame size must be half of FFT size.");
self.far_end_buffer.as_mut_slice().copy_within(self.frame_size.., 0);
self.far_end_buffer
.rows_mut(self.frame_size, self.frame_size)
.copy_from_slice(far_end_frame);
let mut x_t_buffer: Vec<Complex<f32>> = self
.far_end_buffer
.iter()
.map(|&x| Complex::new(x, 0.0))
.collect();
self.fft.process(&mut x_t_buffer);
let x_f = DVector::from_vec(x_t_buffer);
for i in 0..self.fft_size {
let power = x_f[i].norm_sqr();
self.psd[i] = self.smoothing_factor * self.psd[i] + (1.0 - self.smoothing_factor) * power;
}
let y_f = self.weights.component_mul(&x_f);
let mut y_t_complex = y_f.as_slice().to_vec();
self.ifft.process(&mut y_t_complex);
let fft_size_f32 = self.fft_size as f32;
let y_t: DVector<f32> = DVector::from_iterator(
self.fft_size,
y_t_complex.iter().map(|c| c.re / fft_size_f32),
);
let estimated_echo = y_t.rows(self.frame_size, self.frame_size);
let error_signal: Vec<f32> = mic_frame
.iter()
.zip(estimated_echo.iter())
.map(|(mic, echo)| mic - echo)
.collect();
let mut e_t_buffer = vec![Complex::new(0.0, 0.0); self.fft_size];
for (i, &sample) in error_signal.iter().enumerate() {
e_t_buffer[i + self.frame_size] = Complex::new(sample, 0.0);
}
self.fft.process(&mut e_t_buffer);
let e_f = DVector::from_vec(e_t_buffer);
let mut gradient = x_f.map(|c| c.conj()).component_mul(&e_f);
for i in 0..self.fft_size {
gradient[i] /= self.psd[i] + 1e-10; }
self.weights += &gradient * Complex::new(self.mu, 0.0);
error_signal
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_instance_and_process_frame() {
const FFT_SIZE: usize = 512;
const FRAME_SIZE: usize = FFT_SIZE / 2;
const STEP_SIZE: f32 = 0.5;
let mut aec = FdafAec::new(FFT_SIZE, STEP_SIZE);
let far_end_frame = vec![0.0; FRAME_SIZE];
let mic_frame = vec![0.1; FRAME_SIZE];
let error_signal = aec.process(&far_end_frame, &mic_frame);
assert_eq!(error_signal.len(), FRAME_SIZE);
assert!(error_signal.iter().all(|&x| x.is_finite()), "Output contains NaN or Infinity");
}
#[test]
#[should_panic]
fn test_new_with_non_power_of_two_fft_size() {
FdafAec::new(511, 0.5);
}
#[test]
#[should_panic]
fn test_process_with_wrong_frame_size() {
let mut aec = FdafAec::new(512, 0.5);
let far_end_frame = vec![0.0; 128];
let mic_frame = vec![0.0; 256];
aec.process(&far_end_frame, &mic_frame);
}
}