yum/gpu_fft

A GPU-friendly FFT

git clone https://git.yummers.dev/yum/gpu_fft

yumFormatting3c2c603

master
31.8 KiB908 linesraw
1use num::complex::Complex;
2use num::traits::{Float, FloatConst};
3
4fn is_power_of_k(n: usize, k: usize) -> bool {
5    match n {
6        0 => false,
7        1 => true,
8        _ => n % k == 0 && is_power_of_k(n / k, k),
9    }
10}
11
12fn usize_to_float<T: Float>(value: usize) -> T {
13    num::cast(value).unwrap()
14}
15
16#[rustfmt::skip]
17#[allow(dead_code)]
18pub fn naive_dft<T: Float + FloatConst>(data: &mut [Complex<T>]) {
19    let big_n = data.len();
20    let mut result = vec![Complex::new(T::zero(), T::zero()); big_n];
21    for k in 0..big_n {
22        for n in 0..big_n {
23            let k_t     = usize_to_float::<T>(k);
24            let n_t     = usize_to_float::<T>(n);
25            let big_n_t = usize_to_float::<T>(big_n);
26            let phase = -T::TAU() * k_t * n_t / big_n_t;
27            let factor = Complex::<T>::cis(phase);
28            result[k] = result[k] + data[n] * factor;
29        }
30    }
31    data.copy_from_slice(&result);
32}
33
34// Helper to naive_fft. Takes `data` along with 3 numbers that let us recreate an even-odd subset:
35//  - `start_idx` tells us where the subset begins;
36//  - `big_n` is the number of elements in the subset;
37//  - `stride` is the distance between elements.
38// We also use a double buffer, `scratch`, to avoid clobbering data while merging results.
39#[rustfmt::skip]
40fn _naive_fft<T: Float + FloatConst>(data: &mut [Complex<T>], start_idx: usize, big_n: usize, stride: usize, scratch: &mut [Complex<T>]) {
41    if big_n == 1 {
42        return;
43    }
44    // Compute DFT of even elements.
45    _naive_fft(data, start_idx,        big_n/2, stride*2, scratch);
46    // Odd elements.
47    _naive_fft(data, start_idx+stride, big_n/2, stride*2, scratch);
48    for k in 0..(big_n/2) {
49        let p = data[start_idx + 2 * k       * stride];
50        let q = data[start_idx + (2 * k + 1) * stride];
51        let k_t     = usize_to_float::<T>(k);
52        let big_n_t = usize_to_float::<T>(big_n);
53        let phase = -T::TAU() * k_t / big_n_t;
54        let factor = Complex::<T>::cis(phase);
55        scratch[start_idx + k               * stride] = p + q * factor;
56        scratch[start_idx + (k + big_n / 2) * stride] = p - q * factor;
57    }
58    data.copy_from_slice(scratch);
59}
60
61// Naive implementation of Cooley-Tukey FFT. Modifies `data`in place. Panics if data.len() is not a power of two.
62#[allow(dead_code)]
63pub fn naive_fft<T: Float + FloatConst>(data: &mut [Complex<T>]) {
64    assert!(is_power_of_k(data.len(), 2));
65    let mut scratch = Vec::from(data.as_ref());
66    _naive_fft(data, 0, data.len(), 1, &mut scratch);
67}
68
69fn _fft_v1_hoist<T: Float + FloatConst>(
70    data: &mut [Complex<T>],
71    start_idx: usize,
72    big_n: usize,
73    stride: usize,
74    scratch: &mut [Complex<T>],
75    twiddles: &[Complex<T>],
76) {
77    if big_n == 1 {
78        return;
79    }
80    // Compute DFT of even elements.
81    _fft_v1_hoist(data, start_idx, big_n / 2, stride * 2, scratch, twiddles);
82    // Odd elements.
83    _fft_v1_hoist(
84        data,
85        start_idx + stride,
86        big_n / 2,
87        stride * 2,
88        scratch,
89        twiddles,
90    );
91    for k in 0..(big_n / 2) {
92        let p = data[start_idx + 2 * k * stride];
93        let q = data[start_idx + (2 * k + 1) * stride];
94        let factor = twiddles[k * stride];
95        scratch[start_idx + k * stride] = p + q * factor;
96        scratch[start_idx + (k + big_n / 2) * stride] = p - q * factor;
97    }
98    data.copy_from_slice(scratch);
99}
100
101// Modification of fft_naive: hoist out and precompute twiddles.
102pub fn fft_v1_hoist<T: Float + FloatConst>(data: &mut [Complex<T>], twiddles: &[Complex<T>]) {
103    assert!(is_power_of_k(data.len(), 2));
104    let mut scratch = Vec::from(data.as_ref());
105    _fft_v1_hoist(data, 0, data.len(), 1, &mut scratch, &twiddles);
106}
107
108#[inline(always)]
109fn _fft_v2_double_buffer<T: Float + FloatConst>(
110    src: &mut [Complex<T>],
111    dst: &mut [Complex<T>],
112    start_idx: usize,
113    big_n: usize,
114    stride: usize,
115    twiddles: &[Complex<T>],
116) {
117    if big_n == 1 {
118        return;
119    }
120    // Compute DFT of even elements.
121    _fft_v2_double_buffer(dst, src, start_idx, big_n / 2, stride * 2, twiddles);
122    // Odd elements.
123    _fft_v2_double_buffer(
124        dst,
125        src,
126        start_idx + stride,
127        big_n / 2,
128        stride * 2,
129        twiddles,
130    );
131    for k in 0..(big_n / 2) {
132        let p = src[start_idx + 2 * k * stride];
133        let q = src[start_idx + (2 * k + 1) * stride];
134        let factor = twiddles[k * stride];
135        dst[start_idx + k * stride] = p + q * factor;
136        dst[start_idx + (k + big_n / 2) * stride] = p - q * factor;
137    }
138}
139
140pub fn fft_v2_double_buffer<T: Float + FloatConst>(
141    src: &mut [Complex<T>],
142    dst: &mut [Complex<T>],
143    twiddles: &[Complex<T>],
144) {
145    assert!(is_power_of_k(src.len(), 2));
146    dst.copy_from_slice(src);
147    // Switching `src` and `dst` means that at the end, the result is in `src` - which is actually
148    // what we want! We will be hiding `dst` and `twiddles` in a struct later on :)
149    _fft_v2_double_buffer(dst, src, 0, src.len(), 1, twiddles);
150}
151
152// Evaluates the base-`k` logarithm of `n`.
153// Precondition: `n` is a power of `k`.
154const fn log_k_of<const K: usize>(mut n: usize) -> usize {
155    let mut res = 0;
156    while n > 1 {
157        n /= K;
158        res += 1;
159    }
160    res
161}
162
163pub fn fft_v3_iterative<T: Float + FloatConst>(
164    src: &mut [Complex<T>],
165    dst: &mut [Complex<T>],
166    twiddles: &[Complex<T>],
167) {
168    assert!(is_power_of_k(src.len(), 2));
169    let n_iter = log_k_of::<2>(src.len());
170
171    dst.copy_from_slice(src);
172
173    let (mut input, mut output) = if n_iter % 2 == 0 {
174        (dst, src)
175    } else {
176        (src, dst)
177    };
178    let mut stride = input.len();
179    let mut big_n = 1;
180    for _ in 0..n_iter {
181        stride /= 2;
182        big_n *= 2;
183        std::mem::swap(&mut input, &mut output);
184
185        for start_idx in 0..stride {
186            for k in 0..big_n / 2 {
187                // Get odd and even elements.
188                let p = input[start_idx + 2 * k * stride];
189                let q = input[start_idx + (2 * k + 1) * stride];
190                // Combine.
191                let factor = twiddles[k * stride];
192                output[start_idx + k * stride] = p + q * factor;
193                output[start_idx + (k + big_n / 2) * stride] = p - q * factor;
194            }
195        }
196    }
197}
198
199#[inline(always)]
200fn mul_ni<T: Float + FloatConst>(x: Complex<T>) -> Complex<T> {
201    Complex::new(x.im, -x.re)
202}
203
204fn fft_butterfly_radix_4<T: Float + FloatConst>(
205    input: &mut [Complex<T>],
206    output: &mut [Complex<T>],
207    stride: usize,
208    big_n: usize,
209    twiddles: &[Complex<T>],
210) {
211    for start_idx in 0..stride {
212        for k in 0..big_n / 4 {
213            // Collect inputs.
214            let i0 = input[start_idx + 4 * k * stride];
215            let i1 = input[start_idx + (4 * k + 1) * stride];
216            let i2 = input[start_idx + (4 * k + 2) * stride];
217            let i3 = input[start_idx + (4 * k + 3) * stride];
218            // Collect relevant twiddles.
219            let ot1 = twiddles[1 * k * stride];
220            let ot2 = twiddles[2 * k * stride];
221            let ot3 = twiddles[3 * k * stride];
222
223            let a = i0;
224            let b = ot1 * i1;
225            let c = ot2 * i2;
226            let d = ot3 * i3;
227
228            // To derive this, write the output assignments in terms of
229            // a/b/c/d, then factor out!
230            let ac_sum = a + c;
231            let ac_diff = a - c;
232            let bd_sum = b + d;
233            let bd_diff_ni = mul_ni(b - d);
234
235            output[start_idx + k * stride] = ac_sum + bd_sum;
236            output[start_idx + (k + big_n / 4) * stride] = ac_diff + bd_diff_ni;
237            output[start_idx + (k + big_n / 2) * stride] = ac_sum - bd_sum;
238            output[start_idx + (k + 3 * big_n / 4) * stride] = ac_diff - bd_diff_ni;
239        }
240    }
241}
242
243fn fft_butterfly_radix_4_s0<T: Float + FloatConst>(
244    input: &mut [Complex<T>],
245    output: &mut [Complex<T>],
246) {
247    let stride = input.len() / 4;
248    let big_n = 4;
249
250    for start_idx in 0..stride {
251        for k in 0..big_n / 4 {
252            // Collect inputs.
253            let i0 = input[start_idx + 4 * k * stride];
254            let i1 = input[start_idx + (4 * k + 1) * stride];
255            let i2 = input[start_idx + (4 * k + 2) * stride];
256            let i3 = input[start_idx + (4 * k + 3) * stride];
257
258            let a = i0;
259            let b = i1;
260            let c = i2;
261            let d = i3;
262
263            // To derive this, write the output assignments in terms of
264            // a/b/c/d, then factor out!
265            let ac_sum = a + c;
266            let ac_diff = a - c;
267            let bd_sum = b + d;
268            let bd_diff_ni = mul_ni(b - d);
269
270            output[start_idx + k * stride] = ac_sum + bd_sum;
271            output[start_idx + (k + big_n / 4) * stride] = ac_diff + bd_diff_ni;
272            output[start_idx + (k + big_n / 2) * stride] = ac_sum - bd_sum;
273            output[start_idx + (k + 3 * big_n / 4) * stride] = ac_diff - bd_diff_ni;
274        }
275    }
276}
277
278pub fn fft_v4_radix_4<T: Float + FloatConst>(
279    src: &mut [Complex<T>],
280    dst: &mut [Complex<T>],
281    twiddles: &[Complex<T>],
282) {
283    assert!(is_power_of_k(src.len(), 4));
284    let n_iter = log_k_of::<4>(src.len());
285
286    dst.copy_from_slice(src);
287
288    let (mut input, mut output) = if n_iter % 2 == 0 {
289        (dst, src)
290    } else {
291        (src, dst)
292    };
293    let big_n = input.len();
294    let mut stride = big_n;
295    let mut big_n = 1;
296    for _ in 0..n_iter {
297        stride /= 4;
298        big_n *= 4;
299        std::mem::swap(&mut input, &mut output);
300
301        fft_butterfly_radix_4(input, output, stride, big_n, twiddles);
302    }
303}
304
305pub fn fft_v5_s0_opt<T: Float + FloatConst>(
306    src: &mut [Complex<T>],
307    dst: &mut [Complex<T>],
308    twiddles: &[Complex<T>],
309) {
310    assert!(is_power_of_k(src.len(), 4));
311    let n_iter = log_k_of::<4>(src.len());
312
313    dst.copy_from_slice(src);
314
315    let (mut input, mut output) = if n_iter % 2 == 0 {
316        (dst, src)
317    } else {
318        (src, dst)
319    };
320    let big_n = input.len();
321    let mut stride = big_n;
322    let mut big_n = 1;
323    for stage in 0..n_iter {
324        stride /= 4;
325        big_n *= 4;
326        std::mem::swap(&mut input, &mut output);
327
328        if stage == 0 {
329            fft_butterfly_radix_4_s0(input, output);
330        } else {
331            fft_butterfly_radix_4(input, output, stride, big_n, twiddles);
332        }
333    }
334}
335
336fn fft_butterfly_radix_4_unsafe<T: Float + FloatConst>(
337    input: &mut [Complex<T>],
338    output: &mut [Complex<T>],
339    stride: usize,
340    big_n: usize,
341    twiddles: &[Complex<T>],
342) {
343    let input_ptr = input.as_ptr();
344    let output_ptr = output.as_mut_ptr();
345    for start_idx in 0..stride {
346        for k in 0..big_n / 4 {
347            unsafe {
348                // Collect inputs.
349                let i0 = *input_ptr.add(start_idx + 4 * k * stride);
350                let i1 = *input_ptr.add(start_idx + (4 * k + 1) * stride);
351                let i2 = *input_ptr.add(start_idx + (4 * k + 2) * stride);
352                let i3 = *input_ptr.add(start_idx + (4 * k + 3) * stride);
353                // Collect relevant twiddles.
354                let ot1 = twiddles.get_unchecked(1 * k * stride);
355                let ot2 = twiddles.get_unchecked(2 * k * stride);
356                let ot3 = twiddles.get_unchecked(3 * k * stride);
357
358                let a = i0;
359                let b = ot1 * i1;
360                let c = ot2 * i2;
361                let d = ot3 * i3;
362
363                // To derive this, write the output assignments in terms of
364                // a/b/c/d, then factor out!
365                let ac_sum = a + c;
366                let ac_diff = a - c;
367                let bd_sum = b + d;
368                let bd_diff_ni = mul_ni(b - d);
369
370                *output_ptr.add(start_idx + k * stride) = ac_sum + bd_sum;
371                *output_ptr.add(start_idx + (k + big_n / 4) * stride) = ac_diff + bd_diff_ni;
372                *output_ptr.add(start_idx + (k + big_n / 2) * stride) = ac_sum - bd_sum;
373                *output_ptr.add(start_idx + (k + 3 * big_n / 4) * stride) = ac_diff - bd_diff_ni;
374            }
375        }
376    }
377}
378
379fn fft_butterfly_radix_4_s0_unsafe<T: Float + FloatConst>(
380    input: &mut [Complex<T>],
381    output: &mut [Complex<T>],
382) {
383    let stride = input.len() / 4;
384    let big_n = 4;
385    let input_ptr = input.as_ptr();
386    let output_ptr = output.as_mut_ptr();
387    for start_idx in 0..stride {
388        for k in 0..big_n / 4 {
389            unsafe {
390                // Collect inputs.
391                let i0 = *input_ptr.add(start_idx + 4 * k * stride);
392                let i1 = *input_ptr.add(start_idx + (4 * k + 1) * stride);
393                let i2 = *input_ptr.add(start_idx + (4 * k + 2) * stride);
394                let i3 = *input_ptr.add(start_idx + (4 * k + 3) * stride);
395
396                let a = i0;
397                let b = i1;
398                let c = i2;
399                let d = i3;
400
401                // To derive this, write the output assignments in terms of
402                // a/b/c/d, then factor out!
403                let ac_sum = a + c;
404                let ac_diff = a - c;
405                let bd_sum = b + d;
406                let bd_diff_ni = mul_ni(b - d);
407
408                *output_ptr.add(start_idx + k * stride) = ac_sum + bd_sum;
409                *output_ptr.add(start_idx + (k + big_n / 4) * stride) = ac_diff + bd_diff_ni;
410                *output_ptr.add(start_idx + (k + big_n / 2) * stride) = ac_sum - bd_sum;
411                *output_ptr.add(start_idx + (k + 3 * big_n / 4) * stride) = ac_diff - bd_diff_ni;
412            }
413        }
414    }
415}
416
417pub fn fft_v6_unsafe<T: Float + FloatConst>(
418    src: &mut [Complex<T>],
419    dst: &mut [Complex<T>],
420    twiddles: &[Complex<T>],
421) {
422    assert!(is_power_of_k(src.len(), 4));
423    assert_eq!(src.len(), dst.len());
424    assert_eq!(twiddles.len(), src.len());
425    let n_iter = log_k_of::<4>(src.len());
426
427    dst.copy_from_slice(src);
428
429    let (mut input, mut output) = if n_iter % 2 == 0 {
430        (dst, src)
431    } else {
432        (src, dst)
433    };
434    let big_n = input.len();
435    let mut stride = big_n;
436    let mut big_n = 1;
437    for stage in 0..n_iter {
438        stride /= 4;
439        big_n *= 4;
440        std::mem::swap(&mut input, &mut output);
441
442        if stage == 0 {
443            fft_butterfly_radix_4_s0_unsafe(input, output);
444        } else {
445            fft_butterfly_radix_4_unsafe(input, output, stride, big_n, twiddles);
446        }
447    }
448}
449
450#[inline(always)]
451fn rot_45<T: Float + FloatConst>(c: Complex<T>) -> Complex<T> {
452    let s = T::FRAC_1_SQRT_2();
453    // The standard 2D rotation matrix gives:
454    //  [ cos(pi/4) -sin(pi/4)]   [ s -s ]
455    //  [ sin(pi/4)  cos(pi/4)] = [ s  s ]
456    Complex::<T>::new(c.re - c.im, c.re + c.im) * s
457}
458
459#[inline(always)]
460fn rot_90<T: Float + FloatConst>(c: Complex<T>) -> Complex<T> {
461    // The standard 2D rotation matrix gives:
462    //  [ cos(pi/2) -sin(pi/2)]   [ 0 -1 ]
463    //  [ sin(pi/2)  cos(pi/2)] = [ 1  0 ]
464    Complex::<T>::new(-c.im, c.re)
465}
466
467#[inline(always)]
468fn rot_180<T: Float + FloatConst>(c: Complex<T>) -> Complex<T> {
469    // The standard 2D rotation matrix gives:
470    //  [ cos(pi) -sin(pi)]   [ -1  0 ]
471    //  [ sin(pi)  cos(pi)] = [  0 -1 ]
472    Complex::<T>::new(-c.re, -c.im)
473}
474
475#[inline(always)]
476fn rot_270<T: Float + FloatConst>(c: Complex<T>) -> Complex<T> {
477    // The standard 2D rotation matrix gives:
478    //  [ cos(3pi/2) -sin(3pi/2)]   [  0 1 ]
479    //  [ sin(3pi/2)  cos(3pi/2)] = [ -1 0 ]
480    Complex::<T>::new(c.im, -c.re)
481}
482
483fn fft_butterfly_radix_8_unsafe<T: Float + FloatConst>(
484    input: &mut [Complex<T>],
485    output: &mut [Complex<T>],
486    stride: usize,
487    big_n: usize,
488    twiddles: &[Complex<T>],
489) {
490    let input_ptr = input.as_ptr();
491    let output_ptr = output.as_mut_ptr();
492    for start_idx in 0..stride {
493        for k in 0..big_n / 8 {
494            unsafe {
495                // Collect inputs.
496                let i0 = *input_ptr.add(start_idx + 8 * k * stride);
497                let i1 = *input_ptr.add(start_idx + (8 * k + 1) * stride);
498                let i2 = *input_ptr.add(start_idx + (8 * k + 2) * stride);
499                let i3 = *input_ptr.add(start_idx + (8 * k + 3) * stride);
500                let i4 = *input_ptr.add(start_idx + (8 * k + 4) * stride);
501                let i5 = *input_ptr.add(start_idx + (8 * k + 5) * stride);
502                let i6 = *input_ptr.add(start_idx + (8 * k + 6) * stride);
503                let i7 = *input_ptr.add(start_idx + (8 * k + 7) * stride);
504
505                // Collect relevant twiddles.
506                let ot1 = twiddles.get_unchecked(1 * k * stride);
507                let ot2 = twiddles.get_unchecked(2 * k * stride);
508                let ot3 = twiddles.get_unchecked(3 * k * stride);
509                let ot4 = twiddles.get_unchecked(4 * k * stride);
510                let ot5 = twiddles.get_unchecked(5 * k * stride);
511                let ot6 = twiddles.get_unchecked(6 * k * stride);
512                let ot7 = twiddles.get_unchecked(7 * k * stride);
513
514                let a = i0;
515                let b = ot1 * i1;
516                let c = ot2 * i2;
517                let d = ot3 * i3;
518                let e = ot4 * i4;
519                let f = ot5 * i5;
520                let g = ot6 * i6;
521                let h = ot7 * i7;
522
523                let ae_sum = a + e;
524                let ae_diff = a - e;
525                let bf_sum = b + f;
526                let bf_diff = b - f;
527                let cg_sum = c + g;
528                let cg_diff = c - g;
529                let dh_sum = d + h;
530                let dh_diff = d - h;
531
532                let w00 = ae_sum + cg_sum;
533                let w01 = ae_sum - cg_sum;
534                let w10 = ae_diff + rot_270(cg_diff);
535                let w11 = ae_diff - rot_270(cg_diff);
536                let x00 = bf_sum + dh_sum;
537                let x01 = rot_270(bf_sum) + rot_90(dh_sum);
538                let x10 = rot_45(rot_270(bf_diff) + rot_180(dh_diff));
539                let x11 = rot_45(rot_180(bf_diff) + rot_270(dh_diff));
540
541                *output_ptr.add(start_idx + k * stride) = w00 + x00;
542                *output_ptr.add(start_idx + (k + big_n / 8) * stride) = w10 + x10;
543                *output_ptr.add(start_idx + (k + big_n / 4) * stride) = w01 + x01;
544                *output_ptr.add(start_idx + (k + 3 * big_n / 8) * stride) = w11 + x11;
545                *output_ptr.add(start_idx + (k + big_n / 2) * stride) = w00 - x00;
546                *output_ptr.add(start_idx + (k + 5 * big_n / 8) * stride) = w10 - x10;
547                *output_ptr.add(start_idx + (k + 3 * big_n / 4) * stride) = w01 - x01;
548                *output_ptr.add(start_idx + (k + 7 * big_n / 8) * stride) = w11 - x11;
549            }
550        }
551    }
552}
553
554fn fft_butterfly_radix_8_s0_unsafe<T: Float + FloatConst>(
555    input: &mut [Complex<T>],
556    output: &mut [Complex<T>],
557) {
558    let stride = input.len() / 8;
559    let big_n = 8;
560    let input_ptr = input.as_ptr();
561    let output_ptr = output.as_mut_ptr();
562    for start_idx in 0..stride {
563        for k in 0..big_n / 8 {
564            unsafe {
565                // Collect inputs.
566                let i0 = *input_ptr.add(start_idx + 8 * k * stride);
567                let i1 = *input_ptr.add(start_idx + (8 * k + 1) * stride);
568                let i2 = *input_ptr.add(start_idx + (8 * k + 2) * stride);
569                let i3 = *input_ptr.add(start_idx + (8 * k + 3) * stride);
570                let i4 = *input_ptr.add(start_idx + (8 * k + 4) * stride);
571                let i5 = *input_ptr.add(start_idx + (8 * k + 5) * stride);
572                let i6 = *input_ptr.add(start_idx + (8 * k + 6) * stride);
573                let i7 = *input_ptr.add(start_idx + (8 * k + 7) * stride);
574
575                let a = i0;
576                let b = i1;
577                let c = i2;
578                let d = i3;
579                let e = i4;
580                let f = i5;
581                let g = i6;
582                let h = i7;
583
584                let ae_sum = a + e;
585                let ae_diff = a - e;
586                let bf_sum = b + f;
587                let bf_diff = b - f;
588                let cg_sum = c + g;
589                let cg_diff = c - g;
590                let dh_sum = d + h;
591                let dh_diff = d - h;
592
593                let w00 = ae_sum + cg_sum;
594                let w01 = ae_sum - cg_sum;
595                let w10 = ae_diff + rot_270(cg_diff);
596                let w11 = ae_diff - rot_270(cg_diff);
597                let x00 = bf_sum + dh_sum;
598                let x01 = rot_270(bf_sum) + rot_90(dh_sum);
599                let x10 = rot_45(rot_270(bf_diff) + rot_180(dh_diff));
600                let x11 = rot_45(rot_180(bf_diff) + rot_270(dh_diff));
601
602                *output_ptr.add(start_idx + k * stride) = w00 + x00;
603                *output_ptr.add(start_idx + (k + big_n / 8) * stride) = w10 + x10;
604                *output_ptr.add(start_idx + (k + big_n / 4) * stride) = w01 + x01;
605                *output_ptr.add(start_idx + (k + 3 * big_n / 8) * stride) = w11 + x11;
606                *output_ptr.add(start_idx + (k + big_n / 2) * stride) = w00 - x00;
607                *output_ptr.add(start_idx + (k + 5 * big_n / 8) * stride) = w10 - x10;
608                *output_ptr.add(start_idx + (k + 3 * big_n / 4) * stride) = w01 - x01;
609                *output_ptr.add(start_idx + (k + 7 * big_n / 8) * stride) = w11 - x11;
610            }
611        }
612    }
613}
614
615pub fn fft_v7_radix_8<T: Float + FloatConst>(
616    src: &mut [Complex<T>],
617    dst: &mut [Complex<T>],
618    twiddles: &[Complex<T>],
619) {
620    assert!(is_power_of_k(src.len(), 8));
621    assert_eq!(src.len(), dst.len());
622    assert_eq!(twiddles.len(), src.len());
623    let n_iter = log_k_of::<8>(src.len());
624
625    dst.copy_from_slice(src);
626
627    let (mut input, mut output) = if n_iter % 2 == 0 {
628        (dst, src)
629    } else {
630        (src, dst)
631    };
632    let big_n = input.len();
633    let mut stride = big_n;
634    let mut big_n = 1;
635    for stage in 0..n_iter {
636        stride /= 8;
637        big_n *= 8;
638        std::mem::swap(&mut input, &mut output);
639
640        if stage == 0 {
641            fft_butterfly_radix_8_s0_unsafe(input, output);
642        } else {
643            fft_butterfly_radix_8_unsafe(input, output, stride, big_n, twiddles);
644        }
645    }
646}
647
648// Calculates the "twiddle factors" for an n-element FFT, aka all of the nth roots of unity.
649pub fn precompute_twiddles<T: Float + FloatConst>(n: usize) -> Vec<Complex<T>> {
650    let mut result = vec![Complex::<T>::new(T::zero(), T::zero()); n];
651
652    let n_f64 = usize_to_float::<f64>(n);
653    for i in 0..n {
654        let tw_f64 = Complex::<f64>::cis(-f64::TAU() * usize_to_float::<f64>(i) / (n_f64));
655        result[i] = Complex::new(T::from(tw_f64.re).unwrap(), T::from(tw_f64.im).unwrap());
656    }
657
658    result
659}
660
661// Precompute all twiddles factors for a given radix `r` and input size `n`.
662pub fn precompute_all_twiddles<T: Float + FloatConst>(r: usize, n: usize) -> Vec<Vec<Complex<T>>> {
663    assert!(is_power_of_k(n, r));
664
665    let mut result = Vec::new();
666    let mut n_cur = n;
667    while n_cur > 1 {
668        result.push(precompute_twiddles(n_cur));
669        n_cur /= r;
670    }
671    result
672}
673
674#[cfg(test)]
675mod tests {
676    use rand::RngExt;
677    use rand::SeedableRng;
678    use rand::rngs::StdRng;
679
680    use super::*;
681    use num::complex::Complex;
682    use std::time::Instant;
683
684    // Cast one float type to another, truncating if needed.
685    fn t0_to_t1<T0: Float, T1: Float>(val: T0) -> T1 {
686        return num::cast(val).unwrap();
687    }
688
689    // Casts one Complex<T> to another.
690    fn complex_to_t<T0: Float, T1: Float>(data: &[Complex<T0>]) -> Vec<Complex<T1>> {
691        data.iter()
692            .map(|c| Complex::new(t0_to_t1(c.re), t0_to_t1(c.im)))
693            .collect()
694    }
695
696    fn evaluate_results(
697        result_ref_64: &Vec<Complex<f64>>,
698        result_cur_64: &Vec<Complex<f64>>,
699        duration_dft: std::time::Duration,
700        duration_cur: std::time::Duration,
701        algo_name: &str,
702    ) {
703        assert_eq!(result_ref_64.len(), result_cur_64.len());
704
705        println!("  Algorithm:  {}", algo_name);
706        println!("    Duration:  {:?}", duration_cur);
707        println!(
708            "    Speedup:   {:?}",
709            duration_dft.as_nanos() as f32 / duration_cur.as_nanos() as f32
710        );
711
712        let mut max_err: f64 = 0.0;
713        let mut sum_err: f64 = 0.0;
714        // TODO: median, 90p, 99p
715
716        let result_64: Vec<Complex<f64>> = complex_to_t(&result_cur_64);
717
718        for i in 0..result_ref_64.len() {
719            max_err = max_err.max((result_ref_64[i] - result_64[i]).norm());
720            sum_err += (result_ref_64[i] - result_64[i]).norm();
721        }
722        println!("    Max err:   {}", max_err);
723        println!("    Avg err:   {}", sum_err / (result_ref_64.len() as f64));
724    }
725
726    #[test]
727    fn naive_dft_matches_dft() {
728        println!("Testing naive DFT against naive FFT");
729        {
730            let mut data: Vec<Complex<f64>> = Vec::new();
731            data.resize((2 as usize).pow(12), Complex::ZERO);
732            // Randomize data.
733            let mut rng = StdRng::seed_from_u64(42);
734            for i in 0..data.len() {
735                data[i] = Complex::new(rng.random(), rng.random());
736            }
737
738            let mut result_ref: Vec<rustfft::num_complex::Complex<f64>> = complex_to_t(&data);
739            let mut ref_planner = rustfft::FftPlanner::<f64>::new();
740            let ref_fft = ref_planner.plan_fft_forward(result_ref.len());
741            ref_fft.process(&mut result_ref);
742            let result_ref_64: Vec<Complex<f64>> = complex_to_t(&result_ref);
743
744            println!("Input size {}", data.len());
745
746            let mut result_dft: Vec<Complex<f32>> = complex_to_t(&data);
747            let mut begin_instant = Instant::now();
748            naive_dft(&mut result_dft);
749            let duration_dft = begin_instant.elapsed();
750            let result_dft_64: Vec<Complex<f64>> = complex_to_t(&result_dft);
751            evaluate_results(
752                &result_ref_64,
753                &result_dft_64,
754                duration_dft,
755                duration_dft,
756                "Naive DFT",
757            );
758
759            {
760                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
761                begin_instant = Instant::now();
762                naive_fft(&mut result_fft);
763                let duration_fft = begin_instant.elapsed();
764                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
765                evaluate_results(
766                    &result_ref_64,
767                    &result_fft_64,
768                    duration_dft,
769                    duration_fft,
770                    "Naive FFT",
771                );
772            }
773
774            {
775                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
776                let twiddles = precompute_twiddles(data.len());
777                begin_instant = Instant::now();
778                fft_v1_hoist(&mut result_fft, &twiddles);
779                let duration_fft = begin_instant.elapsed();
780                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
781                evaluate_results(
782                    &result_ref_64,
783                    &result_fft_64,
784                    duration_dft,
785                    duration_fft,
786                    "FFT v1 (hoist twiddles)",
787                );
788            }
789
790            {
791                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
792                let twiddles = precompute_twiddles(data.len());
793                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
794                begin_instant = Instant::now();
795                fft_v2_double_buffer(&mut result_fft, &mut scratch, &twiddles);
796                let duration_fft = begin_instant.elapsed();
797                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
798                evaluate_results(
799                    &result_ref_64,
800                    &result_fft_64,
801                    duration_dft,
802                    duration_fft,
803                    "FFT v2 (double buffer)",
804                );
805            }
806
807            {
808                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
809                let twiddles = precompute_twiddles(data.len());
810                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
811                begin_instant = Instant::now();
812                fft_v3_iterative(&mut result_fft, &mut scratch, &twiddles);
813                let duration_fft = begin_instant.elapsed();
814                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
815                evaluate_results(
816                    &result_ref_64,
817                    &result_fft_64,
818                    duration_dft,
819                    duration_fft,
820                    "FFT v3 (iterative)",
821                );
822            }
823
824            {
825                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
826                let twiddles = precompute_twiddles(data.len());
827                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
828                begin_instant = Instant::now();
829                fft_v4_radix_4(&mut result_fft, &mut scratch, &twiddles);
830                let duration_fft = begin_instant.elapsed();
831                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
832                evaluate_results(
833                    &result_ref_64,
834                    &result_fft_64,
835                    duration_dft,
836                    duration_fft,
837                    "FFT v4 (radix 4)",
838                );
839            }
840
841            {
842                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
843                let twiddles = precompute_twiddles(data.len());
844                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
845                begin_instant = Instant::now();
846                fft_v5_s0_opt(&mut result_fft, &mut scratch, &twiddles);
847                let duration_fft = begin_instant.elapsed();
848                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
849                evaluate_results(
850                    &result_ref_64,
851                    &result_fft_64,
852                    duration_dft,
853                    duration_fft,
854                    "FFT v5 (stage 0 opt)",
855                );
856            }
857
858            {
859                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
860                let twiddles = precompute_twiddles(data.len());
861                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
862                begin_instant = Instant::now();
863                fft_v6_unsafe(&mut result_fft, &mut scratch, &twiddles);
864                let duration_fft = begin_instant.elapsed();
865                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
866                evaluate_results(
867                    &result_ref_64,
868                    &result_fft_64,
869                    duration_dft,
870                    duration_fft,
871                    "FFT v6 (unsafe)",
872                );
873            }
874
875            {
876                let mut result_fft: Vec<Complex<f32>> = complex_to_t(&data);
877                let twiddles = precompute_twiddles(data.len());
878                let mut scratch: Vec<Complex<f32>> = vec![Complex::new(0.0, 0.0); data.len()];
879                begin_instant = Instant::now();
880                fft_v7_radix_8(&mut result_fft, &mut scratch, &twiddles);
881                let duration_fft = begin_instant.elapsed();
882                let result_fft_64: Vec<Complex<f64>> = complex_to_t(&result_fft);
883                evaluate_results(
884                    &result_ref_64,
885                    &result_fft_64,
886                    duration_dft,
887                    duration_fft,
888                    "FFT v7 (radix-8)",
889                );
890            }
891
892            let mut result_rfft: Vec<rustfft::num_complex::Complex<f32>> = complex_to_t(&data);
893            let mut rfft_planner = rustfft::FftPlanner::<f32>::new();
894            let rfft = rfft_planner.plan_fft_forward(result_ref.len());
895            begin_instant = Instant::now();
896            rfft.process(&mut result_rfft);
897            let duration_rfft = begin_instant.elapsed();
898            let result_rfft_64: Vec<Complex<f64>> = complex_to_t(&result_rfft);
899            evaluate_results(
900                &result_ref_64,
901                &result_rfft_64,
902                duration_dft,
903                duration_rfft,
904                "RFFT",
905            );
906        }
907    }
908}