Skip to main content

rill_fft/
overlap_add.rs

1// rill-fft/src/overlap_add.rs
2//! Overlap-add convolution using real FFT.
3//!
4//! Efficient for medium-length impulse responses (up to ~16384 samples).
5//! For very long IRs, use `PartitionedConvolver`.
6
7use num_complex::Complex;
8use rill_core::Transcendental;
9use rill_core_dsp::complex_mat::{mul_complex, soa_from_interleaved};
10
11use crate::real_fft::RealFft;
12
13/// Overlap-add frequency-domain convolver.
14///
15/// Uses a real FFT for efficient frequency-domain multiplication.
16/// All scratch buffers are pre-allocated at construction time.
17///
18/// # Parameters
19///
20/// - `T` — sample type (`f32` or `f64`)
21/// - `BUF_SIZE` — processing block size in samples
22pub struct OverlapAddConvolver<T: Transcendental, const BUF_SIZE: usize> {
23    fft_size: usize,
24    fft: RealFft<T>,
25    ir_spectrum: Vec<Complex<T>>,
26    input_buf: Vec<T>,
27    fft_in: Vec<T>,
28    fft_out: Vec<Complex<T>>,
29    product: Vec<Complex<T>>,
30    ifft_out: Vec<T>,
31    overlap: Vec<T>,
32}
33
34impl<T: Transcendental, const BUF_SIZE: usize> OverlapAddConvolver<T, BUF_SIZE> {
35    /// Create a new overlap-add convolver.
36    ///
37    /// `ir_len` is the expected length of the impulse response. The FFT size
38    /// is chosen as the next power of two >= `BUF_SIZE + ir_len - 1`.
39    ///
40    /// # Panics
41    ///
42    /// Panics if the resulting FFT size is less than 4.
43    pub fn new(ir_len: usize) -> Self {
44        let fft_size = rill_core::utils::next_power_of_two(BUF_SIZE + ir_len - 1).max(4);
45        assert!(fft_size >= 4, "FFT size must be at least 4");
46
47        let fft = RealFft::new(fft_size);
48        let half_plus_one = fft_size / 2 + 1;
49        let overlap_len = fft_size - BUF_SIZE;
50
51        Self {
52            fft_size,
53            fft,
54            ir_spectrum: vec![Complex::new(T::ZERO, T::ZERO); half_plus_one],
55            input_buf: vec![T::ZERO; BUF_SIZE],
56            fft_in: vec![T::ZERO; fft_size],
57            fft_out: vec![Complex::new(T::ZERO, T::ZERO); half_plus_one],
58            product: vec![Complex::new(T::ZERO, T::ZERO); half_plus_one],
59            ifft_out: vec![T::ZERO; fft_size],
60            overlap: vec![T::ZERO; overlap_len],
61        }
62    }
63
64    /// Set the impulse response.
65    ///
66    /// Computes and stores the FFT of the zero-padded IR.
67    pub fn set_ir(&mut self, ir: &[T]) {
68        let mut padded = vec![T::ZERO; self.fft_size];
69        let len = ir.len().min(self.fft_size);
70        padded[..len].copy_from_slice(&ir[..len]);
71
72        self.fft.forward(&padded, &mut self.ir_spectrum);
73    }
74
75    /// Returns the FFT size.
76    pub fn fft_size(&self) -> usize {
77        self.fft_size
78    }
79
80    /// Process one block of samples.
81    ///
82    /// `input` must have exactly `BUF_SIZE` elements.
83    /// `output` must have exactly `BUF_SIZE` elements.
84    pub fn process(&mut self, input: &[T], output: &mut [T]) {
85        assert_eq!(input.len(), BUF_SIZE, "input must have BUF_SIZE elements");
86        assert_eq!(output.len(), BUF_SIZE, "output must have BUF_SIZE elements");
87
88        self.input_buf.copy_from_slice(input);
89
90        self.fft_in.fill(T::ZERO);
91        self.fft_in[..BUF_SIZE].copy_from_slice(&self.input_buf);
92
93        self.fft.forward(&self.fft_in, &mut self.fft_out);
94
95        // 4‑bin batch complex multiply via ComplexSoa
96        let len = self.fft_out.len();
97        let mut i = 0usize;
98        while i + 3 < len {
99            let s = soa_from_interleaved(&self.ir_spectrum[i..i + 4]);
100            let f = soa_from_interleaved(&self.fft_out[i..i + 4]);
101            let prod = s * f;
102            let c = prod.to_complexes();
103            self.product[i] = Complex::new(c[0].0, c[0].1);
104            self.product[i + 1] = Complex::new(c[1].0, c[1].1);
105            self.product[i + 2] = Complex::new(c[2].0, c[2].1);
106            self.product[i + 3] = Complex::new(c[3].0, c[3].1);
107            i += 4;
108        }
109        // Scalar remainder
110        while i < len {
111            self.product[i] = mul_complex(self.ir_spectrum[i], self.fft_out[i]);
112            i += 1;
113        }
114
115        self.fft.inverse(&self.product, &mut self.ifft_out);
116
117        for (out, (ifft_val, overlap_val)) in output
118            .iter_mut()
119            .zip(self.ifft_out.iter().zip(self.overlap.iter()))
120        {
121            *out = *ifft_val + *overlap_val;
122        }
123
124        let overlap_len = self.fft_size - BUF_SIZE;
125        for i in 0..overlap_len {
126            self.overlap[i] = self.ifft_out[BUF_SIZE + i];
127        }
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[test]
136    fn test_unit_impulse_is_passthrough() {
137        let mut conv = OverlapAddConvolver::<f32, 8>::new(4);
138        conv.set_ir(&[1.0, 0.0, 0.0, 0.0]);
139
140        let input = [0.5f32, 0.3, -0.2, 0.8, 0.1, -0.5, 0.4, 0.0];
141        let mut output = [0.0f32; 8];
142        conv.process(&input, &mut output);
143
144        for (i, o) in input.iter().zip(output.iter()) {
145            assert!((i - o).abs() < 1e-3, "expected {i}, got {o}");
146        }
147    }
148
149    #[test]
150    fn test_delayed_impulse_is_delay() {
151        let mut conv = OverlapAddConvolver::<f32, 8>::new(4);
152        conv.set_ir(&[0.0, 0.0, 1.0, 0.0]);
153
154        let input = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
155        let mut output = [0.0f32; 8];
156        conv.process(&input, &mut output);
157
158        assert!((output[0] - 0.0).abs() < 1e-3);
159        assert!((output[1] - 0.1).abs() < 0.5);
160        assert!((output[2] - 1.0).abs() < 0.5);
161        assert!((output[3] - 2.0).abs() < 0.5);
162    }
163
164    #[test]
165    fn test_roundtrip_with_direct_conv() {
166        let ir = [0.3f32, 0.5, 0.2, 0.1];
167
168        let mut ola = OverlapAddConvolver::<f32, 8>::new(ir.len());
169        ola.set_ir(&ir);
170
171        let input = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
172        let mut ola_out = [0.0f32; 8];
173        ola.process(&input, &mut ola_out);
174
175        // Compute reference via direct convolution
176        let mut ref_out = [0.0f32; 8];
177        for n in 0..8 {
178            let mut acc = 0.0;
179            for k in 0..ir.len() {
180                if k <= n {
181                    acc += ir[k] * input[n - k];
182                }
183            }
184            ref_out[n] = acc;
185        }
186
187        for (o, r) in ola_out.iter().zip(ref_out.iter()) {
188            assert!((o - r).abs() < 1e-3, "OLA: {o}, ref: {r}");
189        }
190    }
191
192    #[test]
193    fn test_roundtrip_two_blocks() {
194        let ir = [0.3f32, 0.5, 0.2, 0.1];
195
196        let mut conv = OverlapAddConvolver::<f32, 4>::new(ir.len());
197        conv.set_ir(&ir);
198
199        let block1 = [1.0f32, 2.0, 3.0, 4.0];
200        let block2 = [5.0f32, 6.0, 7.0, 8.0];
201
202        let mut out1 = [0.0f32; 4];
203        let mut out2 = [0.0f32; 4];
204
205        conv.process(&block1, &mut out1);
206        conv.process(&block2, &mut out2);
207
208        let full_input = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
209        let mut ref_out = [0.0f32; 8];
210        for n in 0..8 {
211            let mut acc = 0.0;
212            for k in 0..ir.len() {
213                if k <= n {
214                    acc += ir[k] * full_input[n - k];
215                }
216            }
217            ref_out[n] = acc;
218        }
219
220        for (i, (o, r)) in out1
221            .iter()
222            .chain(out2.iter())
223            .zip(ref_out.iter())
224            .enumerate()
225        {
226            assert!((o - r).abs() < 1e-3, "idx {i}: OLA: {o}, ref: {r}");
227        }
228    }
229}