Skip to main content

competitive/math/
fast_fourier_transform.rs

1use super::{AssociatedValue, Complex, ConvolveSteps, One, Zero};
2#[cfg(target_arch = "x86_64")]
3use super::{SimdBackend, advise_huge_pages, simd_backend};
4
5pub enum ConvolveRealFft {}
6
7pub enum RotateCache {}
8impl RotateCache {
9    pub fn ensure(n: usize) {
10        assert_eq!(n.count_ones(), 1, "call with power of two but {}", n);
11        Self::modify(|cache| {
12            let mut m = cache.len();
13            assert!(
14                m.count_ones() <= 1,
15                "length might be power of two but {}",
16                m
17            );
18            if m >= n {
19                return;
20            }
21            cache.reserve_exact(n - m);
22            if cache.is_empty() {
23                cache.push(Complex::one());
24                m += 1;
25            }
26            while m < n {
27                let p = Complex::primitive_nth_root_of_unity(-((m * 4) as f64));
28                for i in 0..m {
29                    cache.push(cache[i] * p);
30                }
31                m <<= 1;
32            }
33            assert_eq!(cache.len(), n);
34        });
35    }
36}
37crate::impl_assoc_value!(RotateCache, Vec<Complex<f64>>, vec![Complex::one()]);
38
39#[cfg(target_arch = "x86_64")]
40pub mod simd {
41    // These primitives are called only after AVX2 and FMA have been detected.
42    #![allow(clippy::missing_safety_doc, unsafe_op_in_unsafe_fn)]
43
44    use super::{AssociatedValue, Complex, RotateCache, advise_huge_pages};
45    use std::arch::x86_64::*;
46
47    #[derive(Clone, Copy, Default)]
48    #[repr(C, align(32))]
49    pub struct Complex4 {
50        pub re: [f64; 4],
51        pub im: [f64; 4],
52    }
53
54    #[target_feature(enable = "avx2,fma")]
55    #[inline]
56    pub unsafe fn load4(value: &Complex4) -> (__m256d, __m256d) {
57        (
58            _mm256_load_pd(value.re.as_ptr()),
59            _mm256_load_pd(value.im.as_ptr()),
60        )
61    }
62
63    #[target_feature(enable = "avx2,fma")]
64    #[inline]
65    pub unsafe fn store4(value: &mut Complex4, re: __m256d, im: __m256d) {
66        _mm256_store_pd(value.re.as_mut_ptr(), re);
67        _mm256_store_pd(value.im.as_mut_ptr(), im);
68    }
69
70    #[target_feature(enable = "avx2,fma")]
71    #[inline]
72    pub unsafe fn mul4(ar: __m256d, ai: __m256d, br: __m256d, bi: __m256d) -> (__m256d, __m256d) {
73        (
74            _mm256_fmsub_pd(ar, br, _mm256_mul_pd(ai, bi)),
75            _mm256_fmadd_pd(ai, br, _mm256_mul_pd(ar, bi)),
76        )
77    }
78
79    #[target_feature(enable = "avx2,fma")]
80    #[inline]
81    pub unsafe fn multiply_accumulate4(
82        rr: &mut __m256d,
83        ri: &mut __m256d,
84        ar: __m256d,
85        ai: __m256d,
86        br: __m256d,
87        bi: __m256d,
88    ) {
89        *rr = _mm256_fmadd_pd(ar, br, *rr);
90        *rr = _mm256_fnmadd_pd(ai, bi, *rr);
91        *ri = _mm256_fmadd_pd(ai, br, *ri);
92        *ri = _mm256_fmadd_pd(ar, bi, *ri);
93    }
94
95    #[inline]
96    pub fn eval_twiddle(cache: &[Complex<f64>], step: usize, n: usize, k: usize) -> Complex<f64> {
97        let k = step * k;
98        let w = cache[(k >> 2) << 1].conjugate();
99        let w = match k & 3 {
100            0 => w,
101            1 => Complex::new(-w.re, -w.im),
102            2 => Complex::new(-w.im, w.re),
103            _ => Complex::new(w.im, -w.re),
104        };
105        cache[step * n].conjugate() * w
106    }
107
108    #[target_feature(enable = "avx2,fma")]
109    pub unsafe fn fft_soa(a: &mut [Complex4]) {
110        let n = a.len() * 4;
111        RotateCache::ensure(n / 2);
112        RotateCache::with(|cache| {
113            let parity = n.trailing_zeros() & 1;
114            for leaf in (0..n).step_by(16) {
115                let mut level = (n + leaf).trailing_zeros();
116                level -= u32::from(level & 1 != parity);
117                while level >= 4 {
118                    let len = 1usize << level;
119                    let q = leaf >> level;
120                    let width = len / 16;
121                    let start = q * width * 4;
122                    let (a, rest) = a[start..start + width * 4].split_at_mut(width);
123                    let (b, rest) = rest.split_at_mut(width);
124                    let (c, d) = rest.split_at_mut(width);
125                    let w1 = eval_twiddle(cache, 4, n >> level, q);
126                    let w2 = w1 * w1;
127                    let w3 = w1 * w2;
128                    let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
129                    let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
130                    let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
131                    for i in 0..width {
132                        let (ar, ai) = load4(&a[i]);
133                        let (br, bi) = load4(&b[i]);
134                        let (cr, ci) = load4(&c[i]);
135                        let (dr, di) = load4(&d[i]);
136                        let (br, bi) = mul4(br, bi, w1r, w1i);
137                        let (cr, ci) = mul4(cr, ci, w2r, w2i);
138                        let (dr, di) = mul4(dr, di, w3r, w3i);
139                        let acr = _mm256_add_pd(ar, cr);
140                        let aci = _mm256_add_pd(ai, ci);
141                        let bdr = _mm256_add_pd(br, dr);
142                        let bdi = _mm256_add_pd(bi, di);
143                        let acd_r = _mm256_sub_pd(ar, cr);
144                        let acd_i = _mm256_sub_pd(ai, ci);
145                        let bdd_r = _mm256_sub_pd(br, dr);
146                        let bdd_i = _mm256_sub_pd(bi, di);
147                        store4(&mut a[i], _mm256_add_pd(acr, bdr), _mm256_add_pd(aci, bdi));
148                        store4(&mut b[i], _mm256_sub_pd(acr, bdr), _mm256_sub_pd(aci, bdi));
149                        store4(
150                            &mut c[i],
151                            _mm256_sub_pd(acd_r, bdd_i),
152                            _mm256_add_pd(acd_i, bdd_r),
153                        );
154                        store4(
155                            &mut d[i],
156                            _mm256_add_pd(acd_r, bdd_i),
157                            _mm256_sub_pd(acd_i, bdd_r),
158                        );
159                    }
160                    level -= 2;
161                }
162            }
163            if parity != 0 {
164                let blocks = n / 8;
165                for k in 0..blocks {
166                    let w = eval_twiddle(cache, 2, blocks, k);
167                    let wr = _mm256_set1_pd(w.re);
168                    let wi = _mm256_set1_pd(w.im);
169                    let (ar, ai) = load4(&a[k * 2]);
170                    let (br, bi) = load4(&a[k * 2 + 1]);
171                    let (br, bi) = mul4(br, bi, wr, wi);
172                    store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
173                    store4(
174                        &mut a[k * 2 + 1],
175                        _mm256_sub_pd(ar, br),
176                        _mm256_sub_pd(ai, bi),
177                    );
178                }
179            }
180        });
181    }
182
183    #[target_feature(enable = "avx2,fma")]
184    pub unsafe fn ifft_soa(a: &mut [Complex4]) {
185        let n = a.len() * 4;
186        RotateCache::ensure(n / 2);
187        RotateCache::with(|cache| {
188            let parity = n.trailing_zeros() & 1;
189            if parity != 0 {
190                let blocks = n / 8;
191                for k in 0..blocks {
192                    let w = eval_twiddle(cache, 2, blocks, k).conjugate();
193                    let wr = _mm256_set1_pd(w.re);
194                    let wi = _mm256_set1_pd(w.im);
195                    let (ar, ai) = load4(&a[k * 2]);
196                    let (br, bi) = load4(&a[k * 2 + 1]);
197                    store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
198                    let (br, bi) = mul4(_mm256_sub_pd(ar, br), _mm256_sub_pd(ai, bi), wr, wi);
199                    store4(&mut a[k * 2 + 1], br, bi);
200                }
201            }
202            for leaf in (12..n).step_by(16) {
203                let max_level = (leaf + 3).trailing_ones();
204                let mut level = 4 + parity;
205                while level <= max_level {
206                    let len = 1usize << level;
207                    let q = leaf >> level;
208                    let width = len / 16;
209                    let start = q * width * 4;
210                    let (a, rest) = a[start..start + width * 4].split_at_mut(width);
211                    let (b, rest) = rest.split_at_mut(width);
212                    let (c, d) = rest.split_at_mut(width);
213                    let w1 = eval_twiddle(cache, 4, n >> level, q).conjugate();
214                    let w2 = w1 * w1;
215                    let w3 = w1 * w2;
216                    let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
217                    let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
218                    let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
219                    for i in 0..width {
220                        let (ar, ai) = load4(&a[i]);
221                        let (br, bi) = load4(&b[i]);
222                        let (cr, ci) = load4(&c[i]);
223                        let (dr, di) = load4(&d[i]);
224                        let abr = _mm256_add_pd(ar, br);
225                        let abi = _mm256_add_pd(ai, bi);
226                        let cdr = _mm256_add_pd(cr, dr);
227                        let cdi = _mm256_add_pd(ci, di);
228                        let abd_r = _mm256_sub_pd(ar, br);
229                        let abd_i = _mm256_sub_pd(ai, bi);
230                        let cdd_r = _mm256_sub_pd(cr, dr);
231                        let cdd_i = _mm256_sub_pd(ci, di);
232                        store4(&mut a[i], _mm256_add_pd(abr, cdr), _mm256_add_pd(abi, cdi));
233                        let (br, bi) = mul4(
234                            _mm256_add_pd(abd_r, cdd_i),
235                            _mm256_sub_pd(abd_i, cdd_r),
236                            w1r,
237                            w1i,
238                        );
239                        store4(&mut b[i], br, bi);
240                        let (cr, ci) =
241                            mul4(_mm256_sub_pd(abr, cdr), _mm256_sub_pd(abi, cdi), w2r, w2i);
242                        store4(&mut c[i], cr, ci);
243                        let (dr, di) = mul4(
244                            _mm256_sub_pd(abd_r, cdd_i),
245                            _mm256_add_pd(abd_i, cdd_r),
246                            w3r,
247                            w3i,
248                        );
249                        store4(&mut d[i], dr, di);
250                    }
251                    level += 2;
252                }
253            }
254            let scale = _mm256_set1_pd(4.0 / n as f64);
255            for value in a {
256                let (re, im) = load4(value);
257                store4(value, _mm256_mul_pd(re, scale), _mm256_mul_pd(im, scale));
258            }
259        });
260    }
261
262    #[target_feature(enable = "avx2,fma")]
263    unsafe fn dot_one_soa(a: &mut [Complex4], b: &[Complex4]) {
264        let n = a.len() * 4;
265        RotateCache::ensure(n / 2);
266        RotateCache::with(|cache| {
267            for i in 0..a.len() {
268                let (mut br, mut bi) = load4(&b[i]);
269                let mut rr = _mm256_setzero_pd();
270                let mut ri = _mm256_setzero_pd();
271                let w = eval_twiddle(cache, 1, a.len(), i);
272                let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
273                let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
274                for lane in 0..4 {
275                    let ar = _mm256_set1_pd(a[i].re[lane]);
276                    let ai = _mm256_set1_pd(a[i].im[lane]);
277                    multiply_accumulate4(&mut rr, &mut ri, ar, ai, br, bi);
278                    if lane != 3 {
279                        br = _mm256_permute4x64_pd::<0x93>(br);
280                        bi = _mm256_permute4x64_pd::<0x93>(bi);
281                        (br, bi) = mul4(br, bi, wr, wi);
282                    }
283                }
284                store4(&mut a[i], rr, ri);
285            }
286        });
287    }
288
289    #[inline]
290    fn pack_f64(values: impl Iterator<Item = f64>, n: usize) -> Vec<Complex4> {
291        let mut result = Vec::with_capacity(n / 4);
292        advise_huge_pages(&mut result);
293        result.resize(n / 4, Complex4::default());
294        for (i, value) in values.enumerate() {
295            if i < n {
296                result[i >> 2].re[i & 3] = value;
297            } else {
298                result[(i - n) >> 2].im[i & 3] = value;
299            }
300        }
301        result
302    }
303
304    #[target_feature(enable = "avx2,fma")]
305    pub unsafe fn convolve_f64_avx2(
306        a: impl ExactSizeIterator<Item = f64>,
307        b: impl ExactSizeIterator<Item = f64>,
308        range: std::ops::Range<usize>,
309    ) -> Vec<f64> {
310        let n = (range.end.next_power_of_two() / 2).max(4);
311        let mut fa = pack_f64(a, n);
312        let mut fb = pack_f64(b, n);
313        fft_soa(&mut fa);
314        fft_soa(&mut fb);
315        dot_one_soa(&mut fa, &fb);
316        drop(fb);
317        ifft_soa(&mut fa);
318        range
319            .map(|i| {
320                if i < n {
321                    fa[i >> 2].re[i & 3]
322                } else {
323                    fa[(i - n) >> 2].im[i & 3]
324                }
325            })
326            .collect()
327    }
328    #[target_feature(enable = "avx2")]
329    pub unsafe fn convolve_i64_avx2(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
330        super::convolve_i64_naive(a, b, len)
331    }
332    #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
333    pub unsafe fn convolve_i64_avx512(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
334        super::convolve_i64_naive(a, b, len)
335    }
336}
337
338fn bit_reverse<T>(f: &mut [T]) {
339    let mut ip = vec![0u32];
340    let mut k = f.len();
341    let mut m = 1;
342    while 2 * m < k {
343        k /= 2;
344        for j in 0..m {
345            ip.push(ip[j] + k as u32);
346        }
347        m *= 2;
348    }
349    if m == k {
350        for i in 1..m {
351            for j in 0..i {
352                let ji = j + ip[i] as usize;
353                let ij = i + ip[j] as usize;
354                f.swap(ji, ij);
355            }
356        }
357    } else {
358        for i in 1..m {
359            for j in 0..i {
360                let ji = j + ip[i] as usize;
361                let ij = i + ip[j] as usize;
362                f.swap(ji, ij);
363                f.swap(ji + m, ij + m);
364            }
365        }
366    }
367}
368
369fn real_twiddles(n: usize, inverse: bool, mut f: impl FnMut(usize, Complex<f64>)) {
370    const BLOCK: usize = 256;
371    let sign = if inverse { 1.0 } else { -1.0 };
372    let step = Complex::primitive_nth_root_of_unity(sign * n as f64);
373    for start in (1..n / 4).step_by(BLOCK) {
374        let mut w = Complex::polar(1.0, sign * std::f64::consts::TAU * start as f64 / n as f64);
375        for k in start..(start + BLOCK).min(n / 4) {
376            f(k, w);
377            w *= step;
378        }
379    }
380}
381
382pub fn transform_real(t: impl IntoIterator<Item = f64>, len: usize) -> Vec<Complex<f64>> {
383    let n = len.max(4).next_power_of_two();
384    let mut f = vec![Complex::zero(); n / 2];
385    for (i, t) in t.into_iter().enumerate() {
386        if i & 1 == 0 {
387            f[i / 2].re = t;
388        } else {
389            f[i / 2].im = t;
390        }
391    }
392    fft(&mut f);
393    bit_reverse(&mut f);
394    f[0] = Complex::new(f[0].re + f[0].im, f[0].re - f[0].im);
395    f[n / 4] = f[n / 4].conjugate();
396    real_twiddles(n, false, |k, wk| {
397        let c = wk.conjugate().transpose() + 1.;
398        let d = c * (f[k] - f[n / 2 - k].conjugate()) * 0.5;
399        f[k] -= d;
400        f[n / 2 - k] += d.conjugate();
401    });
402    f
403}
404
405pub fn inverse_transform_real(mut f: Vec<Complex<f64>>, len: usize) -> Vec<f64> {
406    let n = len.max(4).next_power_of_two();
407    assert_eq!(f.len(), n / 2);
408    f[0] = Complex::new((f[0].re + f[0].im) * 0.5, (f[0].re - f[0].im) * 0.5);
409    f[n / 4] = f[n / 4].conjugate();
410    real_twiddles(n, true, |k, wk| {
411        let c = wk.transpose().conjugate() + 1.;
412        let d = c * (f[k] - f[n / 2 - k].conjugate()) * 0.5;
413        f[k] -= d;
414        f[n / 2 - k] += d.conjugate();
415    });
416    bit_reverse(&mut f);
417    ifft(&mut f);
418    let inv = 1. / (n / 2) as f64;
419    (0..len)
420        .map(|i| inv * if i & 1 == 0 { f[i / 2].re } else { f[i / 2].im })
421        .collect()
422}
423
424#[inline(always)]
425fn convolve_i64_naive(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
426    let (a, b) = if a.len() < b.len() { (b, a) } else { (a, b) };
427    if b.len() == 1 {
428        return a.into_iter().map(|a| a * b[0]).collect();
429    }
430    let mut c = vec![0; len];
431    for (i, a) in a.chunks(1024).enumerate() {
432        for (j, b) in b.iter().enumerate() {
433            let start = i * 1024 + j;
434            for (c, a) in c[start..start + a.len()].iter_mut().zip(a) {
435                *c += *a * *b;
436            }
437        }
438    }
439    c
440}
441
442impl ConvolveSteps for ConvolveRealFft {
443    type T = Vec<i64>;
444    type F = Vec<Complex<f64>>;
445    fn length(t: &Self::T) -> usize {
446        t.len()
447    }
448    fn transform(t: Self::T, len: usize) -> Self::F {
449        transform_real(t.into_iter().map(|t| t as f64), len)
450    }
451    fn inverse_transform(f: Self::F, len: usize) -> Self::T {
452        inverse_transform_real(f, len)
453            .into_iter()
454            .map(|value| value.round() as i64)
455            .collect()
456    }
457    fn convolve(a: Self::T, b: Self::T) -> Self::T {
458        let len = (a.len() + b.len()).saturating_sub(1);
459        // Keep accumulation overflow-free and exact in the FFT's f64 representation.
460        if (a.len().min(b.len()) <= 32 || {
461            let size = len.next_power_of_two();
462            let log = size.ilog2();
463            let limit = crate::avx_helper!(@dispatch simd_backend, SimdBackend;
464                3 * log + 16, 2 * log + 16, 4 * log + 16
465            );
466            2 * a.len() as u128 * b.len() as u128 <= limit as u128 * size as u128
467        }) && a.iter().map(|x| x.unsigned_abs()).max().unwrap_or(0) as u128
468            * b.iter().map(|x| x.unsigned_abs()).max().unwrap_or(0) as u128
469            <= (1u128 << 53) / a.len().min(b.len()).max(1) as u128
470        {
471            return crate::avx_helper!(@dispatch simd_backend, SimdBackend;
472                unsafe {
473                    if a.len().max(b.len()) < 32 {
474                        simd::convolve_i64_avx2(a, b, len)
475                    } else {
476                        simd::convolve_i64_avx512(a, b, len)
477                    }
478                },
479                unsafe { simd::convolve_i64_avx2(a, b, len) },
480                convolve_i64_naive(a, b, len)
481            );
482        }
483        if !a.is_empty() && !b.is_empty() {
484            crate::avx_helper!(@dispatch_avx2_fma return unsafe {
485                simd::convolve_f64_avx2(
486                    a.into_iter().map(|x| x as f64),
487                    b.into_iter().map(|x| x as f64),
488                    0..len,
489                )
490                .into_iter()
491                .map(|x| x.round() as i64)
492                .collect()
493            }, ());
494        }
495        let mut a = Self::transform(a, len);
496        let b = Self::transform(b, len);
497        Self::multiply(&mut a, &b);
498        Self::inverse_transform(a, len)
499    }
500    fn multiply(f: &mut Self::F, g: &Self::F) {
501        assert_eq!(f.len(), g.len());
502        f[0].re *= g[0].re;
503        f[0].im *= g[0].im;
504        for (f, g) in f.iter_mut().zip(g.iter()).skip(1) {
505            *f *= *g;
506        }
507    }
508}
509
510fn middle_product_f64_scalar(
511    a: impl ExactSizeIterator<Item = f64>,
512    b: impl ExactSizeIterator<Item = f64>,
513) -> Vec<f64> {
514    let a_len = a.len();
515    let b_len = b.len();
516    let len = a_len + b_len - 1;
517    let mut a = transform_real(a, len);
518    let b = transform_real(b, len);
519    ConvolveRealFft::multiply(&mut a, &b);
520    inverse_transform_real(a, len)[b_len - 1..a_len].to_vec()
521}
522
523impl ConvolveRealFft {
524    /// Returns coefficients `b.len() - 1..a.len()` of the convolution of `a` and `b`.
525    /// Panics unless `0 < b.len() <= a.len()`.
526    pub fn middle_product_f64(
527        a: impl ExactSizeIterator<Item = f64>,
528        b: impl ExactSizeIterator<Item = f64>,
529    ) -> Vec<f64> {
530        assert!(0 < b.len() && b.len() <= a.len());
531        crate::avx_helper!(@dispatch_avx2_fma return unsafe {
532            let range = b.len() - 1..a.len();
533            simd::convolve_f64_avx2(a, b, range)
534        }, ());
535        middle_product_f64_scalar(a, b)
536    }
537}
538
539macro_rules! fft_kernel {
540    ($a:expr, $cache:expr, $inverse:expr) => {{
541        let a = $a;
542        let cache = $cache;
543        let n = a.len();
544        if $inverse {
545            let mut v = 1;
546            if n.trailing_zeros() & 1 == 1 {
547                for (a, w) in a.as_chunks_mut::<2>().0.iter_mut().zip(cache) {
548                    let y = (a[0] - a[1]) * w.conjugate();
549                    a[0] += a[1];
550                    a[1] = y;
551                }
552                v = 2;
553            }
554            while v < n {
555                for (q, block) in a.chunks_exact_mut(v * 4).enumerate() {
556                    let (a, rest) = block.split_at_mut(v);
557                    let (b, rest) = rest.split_at_mut(v);
558                    let (c, d) = rest.split_at_mut(v);
559                    let w0 = cache[q].conjugate();
560                    let w1 = cache[q << 1].conjugate();
561                    let w3 = w0 * w1;
562                    for (((a, b), c), d) in a.iter_mut().zip(b).zip(c).zip(d) {
563                        let ac0 = *a + *b;
564                        let ac1 = *c + *d;
565                        let bd0 = *a - *b;
566                        let bd1 = *c - *d;
567                        let bd1 = Complex::new(-bd1.im, bd1.re);
568                        *a = ac0 + ac1;
569                        *b = (bd0 + bd1) * w1;
570                        *c = (ac0 - ac1) * w0;
571                        *d = (bd0 - bd1) * w3;
572                    }
573                }
574                v <<= 2;
575            }
576        } else {
577            let mut v = n / 2;
578            while v >= 2 {
579                let l = v / 2;
580                for (q, block) in a.chunks_exact_mut(l * 4).enumerate() {
581                    let (a, rest) = block.split_at_mut(l);
582                    let (b, rest) = rest.split_at_mut(l);
583                    let (c, d) = rest.split_at_mut(l);
584                    let w0 = cache[q];
585                    let w1 = cache[q << 1];
586                    let w3 = w0 * w1;
587                    for (((a, b), c), d) in a.iter_mut().zip(b).zip(c).zip(d) {
588                        let bv = *b * w1;
589                        let cv = *c * w0;
590                        let dv = *d * w3;
591                        let ac0 = *a + cv;
592                        let ac1 = *a - cv;
593                        let bd0 = bv + dv;
594                        let bd1 = bv - dv;
595                        let bd1 = Complex::new(bd1.im, -bd1.re);
596                        *a = ac0 + bd0;
597                        *b = ac0 - bd0;
598                        *c = ac1 + bd1;
599                        *d = ac1 - bd1;
600                    }
601                }
602                v >>= 2;
603            }
604            if v == 1 {
605                for (a, w) in a.as_chunks_mut::<2>().0.iter_mut().zip(cache) {
606                    let y = a[1] * *w;
607                    a[1] = a[0] - y;
608                    a[0] += y;
609                }
610            }
611        }
612    }};
613}
614
615pub fn fft(a: &mut [Complex<f64>]) {
616    fft_dispatch::<false>(a);
617}
618
619pub fn ifft(a: &mut [Complex<f64>]) {
620    fft_dispatch::<true>(a);
621}
622
623fn fft_dispatch<const INVERSE: bool>(a: &mut [Complex<f64>]) {
624    RotateCache::ensure(a.len() / 2);
625    RotateCache::with(|cache| {
626        #[cfg(target_arch = "x86_64")]
627        if a.len() >= 16 {
628            match simd_backend() {
629                SimdBackend::Avx512 => {
630                    return unsafe { fft_avx512::<INVERSE>(a, cache) };
631                }
632                SimdBackend::Avx2 => return unsafe { fft_avx2::<INVERSE>(a, cache) },
633                SimdBackend::Scalar => {}
634            }
635        }
636        fft_kernel!(a, cache, INVERSE);
637    });
638}
639
640#[cfg(target_arch = "x86_64")]
641#[target_feature(enable = "avx2")]
642unsafe fn fft_avx2<const INVERSE: bool>(a: &mut [Complex<f64>], cache: &[Complex<f64>]) {
643    fft_kernel!(a, cache, INVERSE);
644}
645
646#[cfg(target_arch = "x86_64")]
647#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
648unsafe fn fft_avx512<const INVERSE: bool>(a: &mut [Complex<f64>], cache: &[Complex<f64>]) {
649    fft_kernel!(a, cache, INVERSE);
650}
651
652#[test]
653fn test_convolve_fft() {
654    use crate::{rand, tools::Xorshift};
655    let mut rng = Xorshift::default();
656    for log_n in 1..=16 {
657        let n = 1 << log_n;
658        let input: Vec<_> = (0..n)
659            .map(|_| {
660                Complex::new(
661                    (rng.randf() - 0.5) * 200_000.,
662                    (rng.randf() - 0.5) * 200_000.,
663                )
664            })
665            .collect();
666        let mut actual = input.clone();
667        fft(&mut actual);
668        if n <= 64 {
669            for k in 0..n {
670                let mut expected: Complex<f64> = Complex::zero();
671                for (j, value) in input.iter().enumerate() {
672                    let angle = -std::f64::consts::TAU * (j * k) as f64 / n as f64;
673                    expected += *value * Complex::new(angle.cos(), angle.sin());
674                }
675                let actual = actual[k.reverse_bits() >> (usize::BITS - log_n)];
676                assert!((actual.re - expected.re).abs() < 1e-6);
677                assert!((actual.im - expected.im).abs() < 1e-6);
678            }
679        }
680        ifft(&mut actual);
681        for (actual, expected) in actual.into_iter().zip(input) {
682            assert!((actual.re / n as f64 - expected.re).abs() < 1e-7);
683            assert!((actual.im / n as f64 - expected.im).abs() < 1e-7);
684        }
685    }
686    for log_n in 10..=16 {
687        let size = 1 << log_n;
688        for n in size - 1..=size + 1 {
689            let m = rng.random(1..=size + 1);
690            let a: Vec<i64> = rng.random_iter(-128..=128).take(n).collect();
691            let factor = rng.random(1i64..=256) * if rng.gen_bool(0.5) { 1 } else { -1 };
692            let b: Vec<_> = (0..m)
693                .map(|i| if i & 1 == 0 { factor } else { -factor })
694                .collect();
695            let mut prefix = vec![0; n + 1];
696            for (i, value) in a.iter().enumerate() {
697                prefix[i + 1] = prefix[i] + if i & 1 == 0 { *value } else { -*value };
698            }
699            let expected: Vec<_> = (0..n + m - 1)
700                .map(|i| {
701                    (prefix[(i + 1).min(n)] - prefix[(i + 1).saturating_sub(m)])
702                        * if i & 1 == 0 { factor } else { -factor }
703                })
704                .collect();
705            assert_eq!(ConvolveRealFft::convolve(a.clone(), b.clone()), expected);
706            assert_eq!(ConvolveRealFft::convolve(b, a), expected);
707        }
708    }
709    for m in 1..=32 {
710        for n in [
711            rng.random(1..=64),
712            1023,
713            1024,
714            1025,
715            rng.random(1026..=3073),
716        ] {
717            let factor = rng.random(1i64..=16);
718            let limit = (1i64 << 53) / m as i64 / factor;
719            let a: Vec<_> = (0..n)
720                .map(|_| (limit - rng.random(0i64..=128)) * if rng.gen_bool(0.5) { 1 } else { -1 })
721                .collect();
722            let b: Vec<i64> = rng.random_iter(-factor..=factor).take(m).collect();
723            let mut expected = vec![0; n + m - 1];
724            for (i, a) in a.iter().enumerate() {
725                for (j, b) in b.iter().enumerate() {
726                    expected[i + j] += a * b;
727                }
728            }
729            assert_eq!(ConvolveRealFft::convolve(a.clone(), b.clone()), expected);
730            assert_eq!(ConvolveRealFft::convolve(b, a), expected);
731        }
732    }
733    for n in 0..10 {
734        for m in 0..10 {
735            for rn in 0..2 {
736                for rm in 0..2 {
737                    let n = 2usize.pow(n);
738                    let m = 2usize.pow(m);
739                    let n = n - rng.random(0..n) * rn;
740                    let m = m - rng.random(0..m) * rm;
741                    const A: i64 = 100_000;
742                    rand!(rng, a: [-A..=A; n], b: [-A..=A; m]);
743                    let mut c = vec![0; n + m - 1];
744                    for i in 0..n {
745                        for j in 0..m {
746                            c[i + j] += a[i] * b[j];
747                        }
748                    }
749                    let d = ConvolveRealFft::convolve(a, b);
750                    assert_eq!(c, d);
751                }
752            }
753        }
754    }
755}