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 #![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 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 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}