Skip to main content

convolve_large_ntt

Function convolve_large_ntt 

Source
fn convolve_large_ntt<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 700)
667    fn convolve(mut a: Self::T, mut b: Self::T) -> Self::T {
668        let (threshold, naive_threshold) = (100, 60);
669        #[cfg(target_arch = "x86_64")]
670        let (threshold, naive_threshold) = if use_block_ntt::<M>(64) {
671            (
672                60,
673                if M::RANK >= 13 && a.len().max(b.len()) <= 4096 {
674                    18
675                } else if M::RANK >= 19 && a.len().max(b.len()) <= 262144 {
676                    32
677                } else {
678                    34
679                },
680            )
681        } else {
682            (threshold, naive_threshold)
683        };
684        if Self::length(&a).max(Self::length(&b)) <= threshold {
685            return convolve_karatsuba(&a, &b);
686        }
687        if Self::length(&a).min(Self::length(&b)) <= naive_threshold {
688            return convolve_naive(&a, &b);
689        }
690        let len = (Self::length(&a) + Self::length(&b)).saturating_sub(1);
691        let size = len.max(1).next_power_of_two();
692        let max_size = 1usize << M::RANK;
693        #[cfg(target_arch = "x86_64")]
694        let max_size = if use_block_ntt::<M>(size) {
695            max_size << 3
696        } else {
697            max_size
698        };
699        if size > max_size {
700            return convolve_large_ntt(a, b);
701        }
702        if len <= size / 2 + 2 {
703            let xa = a.pop().unwrap();
704            let xb = b.pop().unwrap();
705            let mut c = vec![MInt::<M>::zero(); len];
706            *c.last_mut().unwrap() = xa * xb;
707            for (a, c) in a.iter().zip(&mut c[b.len()..]) {
708                *c += *a * xb;
709            }
710            for (b, c) in b.iter().zip(&mut c[a.len()..]) {
711                *c += *b * xa;
712            }
713            let d = Self::convolve(a, b);
714            for (d, c) in d.into_iter().zip(&mut c) {
715                *c += d;
716            }
717            return c;
718        }
719        let same = a == b;
720        #[cfg(target_arch = "x86_64")]
721        if use_block_ntt::<M>(size) {
722            a.reserve(size - a.len());
723            b.reserve(size - b.len());
724            advise_huge_pages(&mut a);
725            advise_huge_pages(&mut b);
726            a.resize_with(size, Zero::zero);
727            b.resize_with(size, Zero::zero);
728            unsafe { ntt_simd::convolve_blocks_avx2(&mut a, &mut b, same) };
729            a.truncate(len);
730            return a;
731        }
732        let mut a = Self::transform(a, len);
733        if same {
734            for a in a.iter_mut() {
735                *a *= *a;
736            }
737        } else {
738            let b = Self::transform(b, len);
739            Self::multiply(&mut a, &b);
740        }
741        Self::inverse_transform(a, len)
742    }