Skip to main content

convolve_mint_crt

Function convolve_mint_crt 

Source
fn convolve_mint_crt<M, N1, N2, N3>(
    a: Vec<MInt<M>>,
    b: Vec<MInt<M>>,
) -> Vec<MInt<M>>
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 860)
840    fn convolve(a: Self::T, b: Self::T) -> Self::T {
841        let max_len = Self::length(&a).max(Self::length(&b));
842        let min_len = Self::length(&a).min(Self::length(&b));
843        let (balanced, short) = crate::avx_helper!(@dispatch_avx2_fma (30, 10), (384, 128));
844        if max_len <= balanced || min_len <= short {
845            return convolve_karatsuba(&a, &b);
846        }
847        // Limit coefficient growth to leave headroom for FFT roundoff.
848        let fft_limit = crate::avx_helper!(@dispatch_avx2_fma
849            1usize << ((1u64 << 50) / <M as MIntConvert<u32>>::mod_into() as u64).ilog2().min(20), 0);
850        let convolve = |a: Self::T, b: Self::T| {
851            let fft_len = (a.len() + b.len() - 1).next_power_of_two();
852            if fft_len <= 256 && a.len() * b.len() <= fft_len * 8 {
853                return convolve_karatsuba(&a, &b);
854            }
855            if fft_len <= fft_limit {
856                crate::avx_helper!(@dispatch_avx2_fma return unsafe {
857                    convolve_mint_avx2(a, b)
858                }, ());
859            }
860            convolve_mint_crt::<M, N1, N2, N3>(a, b)
861        };
862        let block_len = min_len.next_power_of_two() * 8 - min_len + 1;
863        let block_len = if min_len <= fft_limit / 2 {
864            block_len.min(fft_limit - min_len + 1)
865        } else {
866            block_len
867        };
868        if max_len <= block_len {
869            return convolve(a, b);
870        }
871        let (a, b) = if a.len() >= b.len() { (a, b) } else { (b, a) };
872        let mut result = vec![MInt::<M>::zero(); a.len() + b.len() - 1];
873        for (i, a) in a.chunks(block_len).enumerate() {
874            let product = convolve(a.to_vec(), b.clone());
875            for (value, product) in result[i * block_len..].iter_mut().zip(product) {
876                *value += product;
877            }
878        }
879        result
880    }