1use 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
13pub 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 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 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 pub fn fft_size(&self) -> usize {
77 self.fft_size
78 }
79
80 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 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 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 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}