1use num_complex::Complex;
15use rill_core::Transcendental;
16
17pub struct ComplexFft<T: Transcendental> {
22 size: usize,
23 bit_reverse: Box<[usize]>,
24 twiddle_cos: Box<[T]>,
25 twiddle_sin: Box<[T]>,
26}
27
28impl<T: Transcendental> ComplexFft<T> {
29 pub fn new(size: usize) -> Self {
35 assert!(
36 size.is_power_of_two(),
37 "FFT size must be a power of two, got {size}"
38 );
39 assert!(size >= 2, "FFT size must be at least 2, got {size}");
40
41 let half_size = size / 2;
42 let log2_size = size.trailing_zeros();
43
44 let mut bit_reverse = vec![0usize; size].into_boxed_slice();
45 for i in 0..size {
46 bit_reverse[i] = i.reverse_bits() >> (usize::BITS - log2_size);
47 }
48
49 let twopi = T::from_f64(2.0 * std::f64::consts::PI);
50 let size_t = T::from_usize(size);
51 let mut twiddle_cos = vec![T::ZERO; half_size].into_boxed_slice();
52 let mut twiddle_sin = vec![T::ZERO; half_size].into_boxed_slice();
53 for k in 0..half_size {
54 let angle = twopi * T::from_usize(k) / size_t;
55 twiddle_cos[k] = angle.cos();
56 twiddle_sin[k] = -angle.sin();
57 }
58
59 Self {
60 size,
61 bit_reverse,
62 twiddle_cos,
63 twiddle_sin,
64 }
65 }
66
67 pub fn size(&self) -> usize {
69 self.size
70 }
71
72 pub fn forward(&self, data: &mut [Complex<T>]) {
78 assert_eq!(data.len(), self.size, "data length must match FFT size");
79 self.bit_reverse_permute(data);
80 self.butterfly(data, false);
81 }
82
83 pub fn inverse(&self, data: &mut [Complex<T>]) {
91 assert_eq!(data.len(), self.size, "data length must match FFT size");
92 self.bit_reverse_permute(data);
93 self.butterfly(data, true);
94 let scale = T::ONE / T::from_usize(self.size);
95 for val in data.iter_mut() {
96 *val = Complex::new(val.re * scale, val.im * scale);
97 }
98 }
99
100 pub fn forward_soa(&self, re: &mut [T], im: &mut [T]) {
108 assert_eq!(re.len(), self.size);
109 assert_eq!(im.len(), self.size);
110 self.bit_reverse_soa(re, im);
111 self.butterfly_soa(re, im, false);
112 }
113
114 pub fn inverse_soa(&self, re: &mut [T], im: &mut [T]) {
118 assert_eq!(re.len(), self.size);
119 assert_eq!(im.len(), self.size);
120 self.bit_reverse_soa(re, im);
121 self.butterfly_soa(re, im, true);
122 let scale = T::ONE / T::from_usize(self.size);
123 for i in 0..self.size {
124 re[i] *= scale;
125 im[i] *= scale;
126 }
127 }
128
129 fn bit_reverse_permute(&self, data: &mut [Complex<T>]) {
130 for i in 0..self.size {
131 let j = self.bit_reverse[i];
132 if i < j {
133 data.swap(i, j);
134 }
135 }
136 }
137
138 fn butterfly(&self, data: &mut [Complex<T>], inverse: bool) {
139 let mut step = 2usize;
140 while step <= self.size {
141 let half_step = step / 2;
142 let step_ratio = self.size / step;
143 let mut block = 0usize;
144 while block < self.size {
145 for pair in 0..half_step {
146 let i = block + pair;
147 let j = i + half_step;
148 let twiddle_idx = pair * step_ratio;
149 let w_cos = self.twiddle_cos[twiddle_idx];
150 let mut w_sin = self.twiddle_sin[twiddle_idx];
151 if inverse {
152 w_sin = -w_sin;
153 }
154 let a = data[i];
155 let b = data[j];
156 let b_re = b.re * w_cos - b.im * w_sin;
157 let b_im = b.im * w_cos + b.re * w_sin;
158 data[i] = Complex::new(a.re + b_re, a.im + b_im);
159 data[j] = Complex::new(a.re - b_re, a.im - b_im);
160 }
161 block += step;
162 }
163 step *= 2;
164 }
165 }
166
167 fn bit_reverse_soa(&self, re: &mut [T], im: &mut [T]) {
168 for i in 0..self.size {
169 let j = self.bit_reverse[i];
170 if i < j {
171 re.swap(i, j);
172 im.swap(i, j);
173 }
174 }
175 }
176
177 fn butterfly_soa(&self, re: &mut [T], im: &mut [T], inverse: bool) {
178 let mut step = 2usize;
179 while step <= self.size {
180 let half_step = step / 2;
181 let step_ratio = self.size / step;
182 let mut block = 0usize;
183 while block < self.size {
184 for pair in 0..half_step {
185 let i = block + pair;
186 let j = i + half_step;
187 let twiddle_idx = pair * step_ratio;
188 let w_cos = self.twiddle_cos[twiddle_idx];
189 let mut w_sin = self.twiddle_sin[twiddle_idx];
190 if inverse {
191 w_sin = -w_sin;
192 }
193 let a_re = re[i];
194 let a_im = im[i];
195 let b_re = re[j];
196 let b_im = im[j];
197 let br = b_re * w_cos - b_im * w_sin;
198 let bi = b_im * w_cos + b_re * w_sin;
199 re[i] = a_re + br;
200 im[i] = a_im + bi;
201 re[j] = a_re - br;
202 im[j] = a_im - bi;
203 }
204 block += step;
205 }
206 step *= 2;
207 }
208 }
209}
210
211#[cfg(test)]
212mod tests {
213 use super::*;
214
215 fn approx_eq(a: Complex<f32>, b: Complex<f32>, eps: f32) -> bool {
216 (a.re - b.re).abs() < eps && (a.im - b.im).abs() < eps
217 }
218
219 #[test]
220 fn test_fft_size_4_forward() {
221 let fft = ComplexFft::<f32>::new(4);
222 let mut data = [
223 Complex::new(1.0, 0.0),
224 Complex::new(1.0, 0.0),
225 Complex::new(1.0, 0.0),
226 Complex::new(1.0, 0.0),
227 ];
228 fft.forward(&mut data);
229 assert!((data[0].re - 4.0).abs() < 1e-4);
231 assert!((data[1].re - 0.0).abs() < 1e-4);
232 assert!((data[2].re - 0.0).abs() < 1e-4);
233 assert!((data[3].re - 0.0).abs() < 1e-4);
234 }
235
236 #[test]
237 fn test_fft_size_4_roundtrip() {
238 let fft = ComplexFft::<f32>::new(4);
239 let original = [
240 Complex::new(1.0, 2.0),
241 Complex::new(3.0, 4.0),
242 Complex::new(5.0, 6.0),
243 Complex::new(7.0, 8.0),
244 ];
245 let mut data = original;
246 fft.forward(&mut data);
247 fft.inverse(&mut data);
248 for i in 0..4 {
249 assert!(approx_eq(data[i], original[i], 1e-4));
250 }
251 }
252
253 #[test]
254 fn test_fft_size_8_roundtrip() {
255 let fft = ComplexFft::<f32>::new(8);
256 let original: Vec<_> = (0..8)
257 .map(|i| Complex::new(i as f32, -(i as f32)))
258 .collect();
259 let mut data = original.clone();
260 fft.forward(&mut data);
261 fft.inverse(&mut data);
262 for (a, b) in data.iter().zip(original.iter()) {
263 assert!(approx_eq(*a, *b, 1e-3));
264 }
265 }
266
267 #[test]
268 fn test_fft_size_16_roundtrip() {
269 let fft = ComplexFft::<f32>::new(16);
270 let original: Vec<_> = (0..16)
271 .map(|i| Complex::new((i as f32).sin(), (i as f32).cos()))
272 .collect();
273 let mut data = original.clone();
274 fft.forward(&mut data);
275 fft.inverse(&mut data);
276 for (a, b) in data.iter().zip(original.iter()) {
277 assert!(approx_eq(*a, *b, 1e-3));
278 }
279 }
280
281 #[test]
282 fn test_fft_impulse_is_constant() {
283 let fft = ComplexFft::<f32>::new(8);
284 let mut data = [Complex::new(0.0, 0.0); 8];
285 data[0] = Complex::new(1.0, 0.0);
286 fft.forward(&mut data);
287 let expected = 1.0;
288 for val in &data {
289 assert!((val.re - expected).abs() < 1e-4);
290 assert!(val.im.abs() < 1e-4);
291 }
292 }
293
294 #[test]
295 fn test_fft_f64_roundtrip() {
296 let fft = ComplexFft::<f64>::new(8);
297 let original: Vec<_> = (0..8)
298 .map(|i| Complex::new((i as f64).sin(), (i as f64).cos()))
299 .collect();
300 let mut data = original.clone();
301 fft.forward(&mut data);
302 fft.inverse(&mut data);
303 for (a, b) in data.iter().zip(original.iter()) {
304 assert!((a.re - b.re).abs() < 1e-10);
305 assert!((a.im - b.im).abs() < 1e-10);
306 }
307 }
308
309 #[test]
310 fn test_fft_size_1024_roundtrip() {
311 let fft = ComplexFft::<f32>::new(1024);
312 let original: Vec<_> = (0..1024)
313 .map(|i| {
314 let x = i as f32 * 0.01;
315 Complex::new(x.sin(), x.cos())
316 })
317 .collect();
318 let mut data = original.clone();
319 fft.forward(&mut data);
320 fft.inverse(&mut data);
321 for (a, b) in data.iter().zip(original.iter()) {
322 assert!(approx_eq(*a, *b, 5e-4));
323 }
324 }
325
326 #[test]
327 fn test_fft_dc_offset() {
328 let fft = ComplexFft::<f32>::new(16);
329 let mut data = [Complex::new(3.0, 0.0); 16];
330 fft.forward(&mut data);
331 assert!((data[0].re - 48.0).abs() < 1e-3);
332 for i in 1..16 {
333 assert!((data[i].re).abs() < 1e-3);
334 assert!((data[i].im).abs() < 1e-3);
335 }
336 }
337
338 #[test]
339 #[should_panic(expected = "power of two")]
340 fn test_fft_non_power_of_two_panics() {
341 ComplexFft::<f32>::new(10);
342 }
343
344 #[test]
345 #[should_panic(expected = "at least 2")]
346 fn test_fft_size_one_panics() {
347 ComplexFft::<f32>::new(1);
348 }
349
350 #[test]
351 fn test_fft_size_2_basic() {
352 let fft = ComplexFft::<f32>::new(2);
353 let mut data = [Complex::new(1.0, 0.0), Complex::new(-1.0, 0.0)];
354 fft.forward(&mut data);
355 assert!((data[0].re - 0.0).abs() < 1e-4);
356 assert!((data[1].re - 2.0).abs() < 1e-4);
357 }
358
359 #[test]
364 fn test_soa_roundtrip() {
365 let fft = ComplexFft::<f32>::new(1024);
366 let mut re: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.01).sin()).collect();
367 let mut im: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.01).cos()).collect();
368 let re_orig = re.clone();
369 let im_orig = im.clone();
370
371 fft.forward_soa(&mut re, &mut im);
372 fft.inverse_soa(&mut re, &mut im);
373
374 for (a, b) in re.iter().zip(re_orig.iter()) {
375 assert!((a - b).abs() < 5e-4);
376 }
377 for (a, b) in im.iter().zip(im_orig.iter()) {
378 assert!((a - b).abs() < 5e-4);
379 }
380 }
381
382 #[test]
383 fn test_soa_f64_roundtrip() {
384 let fft = ComplexFft::<f64>::new(256);
385 let mut re: Vec<f64> = (0..256).map(|i| (i as f64 * 0.01).sin()).collect();
386 let mut im: Vec<f64> = (0..256).map(|i| (i as f64 * 0.01).cos()).collect();
387 let re_orig = re.clone();
388 let im_orig = im.clone();
389
390 fft.forward_soa(&mut re, &mut im);
391 fft.inverse_soa(&mut re, &mut im);
392
393 for (a, b) in re.iter().zip(re_orig.iter()) {
394 assert!((a - b).abs() < 1e-8);
395 }
396 for (a, b) in im.iter().zip(im_orig.iter()) {
397 assert!((a - b).abs() < 1e-8);
398 }
399 }
400
401 #[test]
402 fn test_soa_versus_interleaved() {
403 let fft = ComplexFft::<f32>::new(256);
404 let mut re: Vec<f32> = (0..256).map(|i| (i as f32 * 0.01).sin()).collect();
405 let mut im: Vec<f32> = (0..256).map(|i| (i as f32 * 0.01).cos()).collect();
406 let interleaved: Vec<_> = re
407 .iter()
408 .zip(im.iter())
409 .map(|(&r, &i)| Complex::new(r, i))
410 .collect();
411 let mut inter_ref = interleaved.clone();
412
413 fft.forward(&mut inter_ref);
414 fft.forward_soa(&mut re, &mut im);
415
416 for (idx, (r_val, i_val)) in re.iter().zip(im.iter()).enumerate() {
417 let diff_re = (r_val - inter_ref[idx].re).abs();
418 let diff_im = (i_val - inter_ref[idx].im).abs();
419 assert!(diff_re < 1e-3, "bin {idx}: re diff {diff_re}");
420 assert!(diff_im < 1e-3, "bin {idx}: im diff {diff_im}");
421 }
422 }
423}
424
425#[cfg(feature = "simd")]
430mod simd_soa {
431 use super::*;
432 use rill_core::math::vector::simd::wide::{F32x4, F64x2};
433 use rill_core::math::vector::traits::Vector;
434
435 impl ComplexFft<f32> {
436 pub fn forward_simd(&mut self, re: &mut [f32], im: &mut [f32]) {
441 Self::bit_reverse_soa_simd(&self.bit_reverse, re, im, self.size);
442 Self::butterfly_simd_f32(
443 &self.twiddle_cos,
444 &self.twiddle_sin,
445 self.size,
446 re,
447 im,
448 false,
449 );
450 }
451
452 pub fn inverse_simd(&mut self, re: &mut [f32], im: &mut [f32]) {
454 Self::bit_reverse_soa_simd(&self.bit_reverse, re, im, self.size);
455 Self::butterfly_simd_f32(
456 &self.twiddle_cos,
457 &self.twiddle_sin,
458 self.size,
459 re,
460 im,
461 true,
462 );
463 let scale = 1.0f32 / self.size as f32;
464 for i in 0..self.size {
465 re[i] *= scale;
466 im[i] *= scale;
467 }
468 }
469
470 fn butterfly_simd_f32(
471 tw_cos: &[f32],
472 tw_sin: &[f32],
473 size: usize,
474 re: &mut [f32],
475 im: &mut [f32],
476 inverse: bool,
477 ) {
478 let mut step = 2usize;
479 while step <= size {
480 let half_step = step / 2;
481 let step_ratio = size / step;
482 let mut block = 0usize;
483 while block < size {
484 let mut pair = 0usize;
485 if step_ratio == 1 {
487 while pair + 3 < half_step {
488 let base = block + pair;
489 let base_b = base + half_step;
490 let w_cos = F32x4::load(&tw_cos[pair..]);
491 let mut w_sin = F32x4::load(&tw_sin[pair..]);
492 if inverse {
493 w_sin = -w_sin;
494 }
495 let a_re = F32x4::load(&re[base..]);
496 let a_im = F32x4::load(&im[base..]);
497 let b_re = F32x4::load(&re[base_b..]);
498 let b_im = F32x4::load(&im[base_b..]);
499 let br = b_re * w_cos - b_im * w_sin;
500 let bi = b_im * w_cos + b_re * w_sin;
501 (a_re + br).store(&mut re[base..]);
502 (a_im + bi).store(&mut im[base..]);
503 (a_re - br).store(&mut re[base_b..]);
504 (a_im - bi).store(&mut im[base_b..]);
505 pair += 4;
506 }
507 }
508 while pair < half_step {
510 let i = block + pair;
511 let j = i + half_step;
512 let ti = pair * step_ratio;
513 let wc = tw_cos[ti];
514 let mut ws = tw_sin[ti];
515 if inverse {
516 ws = -ws;
517 }
518 let ar = re[i];
519 let ai = im[i];
520 let br = re[j];
521 let bi = im[j];
522 let br2 = br * wc - bi * ws;
523 let bi2 = bi * wc + br * ws;
524 re[i] = ar + br2;
525 im[i] = ai + bi2;
526 re[j] = ar - br2;
527 im[j] = ai - bi2;
528 pair += 1;
529 }
530 block += step;
531 }
532 step *= 2;
533 }
534 }
535
536 fn bit_reverse_soa_simd(bit_rev: &[usize], re: &mut [f32], im: &mut [f32], _size: usize) {
537 for (i, &j) in bit_rev.iter().enumerate() {
538 if i < j {
539 re.swap(i, j);
540 im.swap(i, j);
541 }
542 }
543 }
544 }
545
546 impl ComplexFft<f64> {
547 pub fn forward_simd(&mut self, re: &mut [f64], im: &mut [f64]) {
549 Self::bit_reverse_soa_simd(&self.bit_reverse, re, im, self.size);
550 Self::butterfly_simd_f64(
551 &self.twiddle_cos,
552 &self.twiddle_sin,
553 self.size,
554 re,
555 im,
556 false,
557 );
558 }
559
560 pub fn inverse_simd(&mut self, re: &mut [f64], im: &mut [f64]) {
562 Self::bit_reverse_soa_simd(&self.bit_reverse, re, im, self.size);
563 Self::butterfly_simd_f64(
564 &self.twiddle_cos,
565 &self.twiddle_sin,
566 self.size,
567 re,
568 im,
569 true,
570 );
571 let scale = 1.0f64 / self.size as f64;
572 for i in 0..self.size {
573 re[i] *= scale;
574 im[i] *= scale;
575 }
576 }
577
578 fn butterfly_simd_f64(
579 tw_cos: &[f64],
580 tw_sin: &[f64],
581 size: usize,
582 re: &mut [f64],
583 im: &mut [f64],
584 inverse: bool,
585 ) {
586 let mut step = 2usize;
587 while step <= size {
588 let half_step = step / 2;
589 let step_ratio = size / step;
590 let mut block = 0usize;
591 while block < size {
592 let mut pair = 0usize;
593 if step_ratio == 1 {
594 while pair + 1 < half_step {
595 let base = block + pair;
596 let base_b = base + half_step;
597 let w_cos = F64x2::load(&tw_cos[pair..]);
598 let mut w_sin = F64x2::load(&tw_sin[pair..]);
599 if inverse {
600 w_sin = -w_sin;
601 }
602 let a_re = F64x2::load(&re[base..]);
603 let a_im = F64x2::load(&im[base..]);
604 let b_re = F64x2::load(&re[base_b..]);
605 let b_im = F64x2::load(&im[base_b..]);
606 let br = b_re * w_cos - b_im * w_sin;
607 let bi = b_im * w_cos + b_re * w_sin;
608 (a_re + br).store(&mut re[base..]);
609 (a_im + bi).store(&mut im[base..]);
610 (a_re - br).store(&mut re[base_b..]);
611 (a_im - bi).store(&mut im[base_b..]);
612 pair += 2;
613 }
614 }
615 while pair < half_step {
616 let i = block + pair;
617 let j = i + half_step;
618 let ti = pair * step_ratio;
619 let wc = tw_cos[ti];
620 let mut ws = tw_sin[ti];
621 if inverse {
622 ws = -ws;
623 }
624 let ar = re[i];
625 let ai = im[i];
626 let br = re[j];
627 let bi = im[j];
628 let br2 = br * wc - bi * ws;
629 let bi2 = bi * wc + br * ws;
630 re[i] = ar + br2;
631 im[i] = ai + bi2;
632 re[j] = ar - br2;
633 im[j] = ai - bi2;
634 pair += 1;
635 }
636 block += step;
637 }
638 step *= 2;
639 }
640 }
641
642 fn bit_reverse_soa_simd(bit_rev: &[usize], re: &mut [f64], im: &mut [f64], _size: usize) {
643 for (i, &j) in bit_rev.iter().enumerate() {
644 if i < j {
645 re.swap(i, j);
646 im.swap(i, j);
647 }
648 }
649 }
650 }
651
652 #[cfg(test)]
653 mod tests {
654 use super::*;
655
656 #[test]
657 fn test_simd_soa_f32_roundtrip() {
658 let mut fft = ComplexFft::<f32>::new(1024);
659 let mut re: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.01).sin()).collect();
660 let mut im: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.01).cos()).collect();
661 let re_orig = re.clone();
662 let im_orig = im.clone();
663
664 fft.forward_simd(&mut re, &mut im);
665 fft.inverse_simd(&mut re, &mut im);
666
667 for (a, b) in re.iter().zip(re_orig.iter()) {
668 assert!((a - b).abs() < 5e-4);
669 }
670 for (a, b) in im.iter().zip(im_orig.iter()) {
671 assert!((a - b).abs() < 5e-4);
672 }
673 }
674
675 #[test]
676 fn test_simd_soa_f64_roundtrip() {
677 let mut fft = ComplexFft::<f64>::new(512);
678 let mut re: Vec<f64> = (0..512).map(|i| (i as f64 * 0.01).sin()).collect();
679 let mut im: Vec<f64> = (0..512).map(|i| (i as f64 * 0.01).cos()).collect();
680 let re_orig = re.clone();
681 let im_orig = im.clone();
682
683 fft.forward_simd(&mut re, &mut im);
684 fft.inverse_simd(&mut re, &mut im);
685
686 for (a, b) in re.iter().zip(re_orig.iter()) {
687 assert!((a - b).abs() < 1e-7);
688 }
689 for (a, b) in im.iter().zip(im_orig.iter()) {
690 assert!((a - b).abs() < 1e-7);
691 }
692 }
693 }
694}