Skip to main content

convolve_i64_naive

Function convolve_i64_naive 

Source
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    }