fn convolve_i64_naive(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64>Examples found in repository?
crates/competitive/src/math/fast_fourier_transform.rs (line 330)
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 }