Skip to main content

RotateCache

Enum RotateCache 

Source
pub enum RotateCache {}

Implementations§

Source§

impl RotateCache

Source

pub fn ensure(n: usize)

Examples found in repository?
crates/competitive/src/math/mint_fft_convolve.rs (line 81)
79unsafe fn dot_soa(a0: &mut [Complex4], a1: &mut [Complex4], b0: &mut [Complex4], b1: &[Complex4]) {
80    let n = a0.len() * 4;
81    RotateCache::ensure(n / 2);
82    RotateCache::with(|cache| {
83        for i in 0..a0.len() {
84            let (mut cr, mut ci) = load4(&b0[i]);
85            let (mut dr, mut di) = load4(&b1[i]);
86            let mut c0r = _mm256_setzero_pd();
87            let mut c0i = _mm256_setzero_pd();
88            let mut c1r = _mm256_setzero_pd();
89            let mut c1i = _mm256_setzero_pd();
90            let mut c2r = _mm256_setzero_pd();
91            let mut c2i = _mm256_setzero_pd();
92            let w = eval_twiddle(cache, 1, a0.len(), i);
93            let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
94            let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
95            for lane in 0..4 {
96                let ar = _mm256_set1_pd(a0[i].re[lane]);
97                let ai = _mm256_set1_pd(a0[i].im[lane]);
98                let br = _mm256_set1_pd(a1[i].re[lane]);
99                let bi = _mm256_set1_pd(a1[i].im[lane]);
100                multiply_accumulate4(&mut c0r, &mut c0i, ar, ai, cr, ci);
101                multiply_accumulate4(&mut c1r, &mut c1i, ar, ai, dr, di);
102                multiply_accumulate4(&mut c1r, &mut c1i, br, bi, cr, ci);
103                multiply_accumulate4(&mut c2r, &mut c2i, br, bi, dr, di);
104                if lane != 3 {
105                    cr = _mm256_permute4x64_pd::<0x93>(cr);
106                    ci = _mm256_permute4x64_pd::<0x93>(ci);
107                    dr = _mm256_permute4x64_pd::<0x93>(dr);
108                    di = _mm256_permute4x64_pd::<0x93>(di);
109                    (cr, ci) = mul4(cr, ci, wr, wi);
110                    (dr, di) = mul4(dr, di, wr, wi);
111                }
112            }
113            store4(&mut a0[i], c0r, c0i);
114            store4(&mut a1[i], c1r, c1i);
115            store4(&mut b0[i], c2r, c2i);
116        }
117    });
118}
119
120#[target_feature(enable = "avx2,fma")]
121unsafe fn split_u64_coefficients(values: &[u64], n: usize) -> [Vec<Complex4>; 5] {
122    let mut result: [Vec<Complex4>; 5] = std::array::from_fn(|_| {
123        let mut part = Vec::with_capacity(n / 4);
124        advise_huge_pages(&mut part);
125        part
126    });
127    for (i, chunk) in values.chunks(4).enumerate() {
128        let mut parts = [Complex4::default(); 5];
129        for (lane, mut value) in chunk.iter().copied().enumerate() {
130            for part in &mut parts {
131                let digit = ((value << 51) as i64) >> 51;
132                value = (value >> 13).wrapping_add(u64::from(digit < 0));
133                part.re[lane] = digit as f64;
134            }
135        }
136        for (result, part) in result.iter_mut().zip(parts) {
137            if i < n / 4 {
138                result.push(part);
139            } else {
140                result[i - n / 4].im = part.re;
141            }
142        }
143    }
144    for part in &mut result {
145        part.resize(n / 4, Complex4::default());
146    }
147    result
148}
149
150#[target_feature(enable = "avx2,fma")]
151unsafe fn dot_u64_soa(a: &mut [Vec<Complex4>; 5], b: &[Vec<Complex4>; 5]) {
152    let n = a[0].len() * 4;
153    RotateCache::ensure(n / 2);
154    RotateCache::with(|cache| {
155        for block in 0..a[0].len() {
156            let mut br = [_mm256_setzero_pd(); 5];
157            let mut bi = br;
158            let mut rr = br;
159            let mut ri = br;
160            for part in 0..5 {
161                (br[part], bi[part]) = load4(&b[part][block]);
162            }
163            let w = eval_twiddle(cache, 1, a[0].len(), block);
164            let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
165            let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
166            for lane in 0..4 {
167                let ar: [__m256d; 5] =
168                    std::array::from_fn(|part| _mm256_set1_pd(a[part][block].re[lane]));
169                let ai: [__m256d; 5] =
170                    std::array::from_fn(|part| _mm256_set1_pd(a[part][block].im[lane]));
171                for part in 0..5 {
172                    for left in 0..=part {
173                        multiply_accumulate4(
174                            &mut rr[part],
175                            &mut ri[part],
176                            ar[left],
177                            ai[left],
178                            br[part - left],
179                            bi[part - left],
180                        );
181                    }
182                }
183                if lane != 3 {
184                    for part in 0..5 {
185                        br[part] = _mm256_permute4x64_pd::<0x93>(br[part]);
186                        bi[part] = _mm256_permute4x64_pd::<0x93>(bi[part]);
187                        (br[part], bi[part]) = mul4(br[part], bi[part], wr, wi);
188                    }
189                }
190            }
191            for part in 0..5 {
192                store4(&mut a[part][block], rr[part], ri[part]);
193            }
194        }
195    });
196}
More examples
Hide additional examples
crates/competitive/src/math/fast_fourier_transform.rs (line 111)
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}

Trait Implementations§

Source§

impl AssociatedValue for RotateCache

Source§

type T = Vec<Complex<f64>>

Type of value.
Source§

unsafe fn __local_key() -> &'static LocalKey<Cell<Self::T>>

Source§

fn get() -> Self::T

Source§

fn set(x: Self::T)

Source§

fn replace(x: Self::T) -> Self::T

Source§

fn with<F, R>(f: F) -> R
where F: FnOnce(&Self::T) -> R,

Source§

fn modify<F, R>(f: F) -> R
where F: FnOnce(&mut Self::T) -> R,

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> ToArrayVecScalar for T

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.