Skip to main content

NttReuse

Trait NttReuse 

Source
pub trait NttReuse: ConvolveSteps {
    const MULTIPLE: bool = true;

    // Required methods
    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F;
    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
    fn multiply_prefix(f: &mut Self::F, g: &Self::F);
    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F);

    // Provided methods
    fn transform_ntt(t: Self::T, len: usize) -> Self::F { ... }
    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T { ... }
    fn max_product_sum_count(_f: &Self::F) -> usize { ... }
    fn power_projection_step(
        p_flat: Self::T,
        q_flat: Self::T,
        n: usize,
        py: usize,
        qy: usize,
    ) -> (Self::T, Self::T) { ... }
}

Provided Associated Constants§

Source

const MULTIPLE: bool = true

Required Methods§

Source

fn ntt_doubling(f: Self::F, monic: bool) -> Self::F

Extends a value produced by transform_ntt to twice its length. If monic, the input represents a monic degree-n polynomial modulo x^n - 1, where n is the transform length.

Source

fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F

Extracts the even coefficients of a(x) * b(-x) in the usual NTT frequency order.

Source

fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F

Extracts the odd coefficients of a(x) * b(-x) in the usual NTT frequency order.

Source

fn multiply_prefix(f: &mut Self::F, g: &Self::F)

Multiplies a usual NTT transform by the corresponding prefix of another one.

Source

fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F)

Adds the pointwise product of two usual NTT transforms to sum.

Provided Methods§

Source

fn transform_ntt(t: Self::T, len: usize) -> Self::F

Transforms coefficients into the usual NTT frequency order.

Examples found in repository?
crates/competitive/src/math/formal_power_series/berlekamp_massey.rs (line 345)
336fn reduced_transform<T, C>(fps: &FormalPowerSeries<T, C>, length: usize) -> C::F
337where
338    T: FormalPowerSeriesCoefficient,
339    C: NttReuse<T = Vec<T>>,
340{
341    let mut coefficients = vec![T::zero(); length];
342    for (i, value) in fps.iter().enumerate() {
343        coefficients[i & (length - 1)] += value;
344    }
345    C::transform_ntt(coefficients, length)
346}
347
348fn transform_window<T, C>(fps: &FormalPowerSeries<T, C>, end: isize, length: usize) -> C::F
349where
350    T: FormalPowerSeriesCoefficient,
351    C: NttReuse<T = Vec<T>>,
352{
353    let start = end - length as isize;
354    let coefficients = (start..end).map(|index| coefficient(fps, index)).collect();
355    C::transform_ntt(coefficients, length)
356}
More examples
Hide additional examples
crates/competitive/src/math/number_theoretic_transform.rs (line 593)
581fn convolve_large_ntt<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
582where
583    M: Montgomery32NttModulus,
584{
585    let len = a.len() + b.len() - 1;
586    let ntt_len = 1usize << M::RANK;
587    let block_len = ntt_len / 2;
588    let same = a == b;
589    let transform = |a: &[MInt<M>]| {
590        let mut f = Vec::with_capacity(ntt_len);
591        advise_huge_pages(&mut f);
592        f.extend_from_slice(a);
593        Convolve::<M>::transform_ntt(f, ntt_len)
594    };
595    let fa: Vec<_> = a.chunks(block_len).map(transform).collect();
596    let fb: Option<Vec<_>> = if same {
597        None
598    } else {
599        Some(b.chunks(block_len).map(transform).collect())
600    };
601    let b_blocks = fb.as_ref().map_or(fa.len(), Vec::len);
602    let mut result = vec![MInt::<M>::zero(); len];
603    for diagonal in 0..fa.len() + b_blocks - 1 {
604        let mut spectrum = vec![MInt::<M>::zero(); ntt_len];
605        let start = diagonal.saturating_sub(b_blocks - 1);
606        for i in start..=diagonal.min(fa.len() - 1) {
607            let j = diagonal - i;
608            let g = if let Some(fb) = &fb { &fb[j] } else { &fa[j] };
609            pointwise_multiply_add(&mut spectrum, &fa[i], g);
610        }
611        spectrum = Convolve::<M>::inverse_transform_ntt(spectrum, ntt_len);
612        let offset = diagonal * block_len;
613        for (result, value) in result[offset..].iter_mut().zip(spectrum) {
614            *result += value;
615        }
616    }
617    result
618}
619
620impl<M> ConvolveSteps for Convolve<M>
621where
622    M: Montgomery32NttModulus,
623{
624    const CYCLIC: bool = true;
625
626    type T = Vec<MInt<M>>;
627    type F = Vec<MInt<M>>;
628    fn length(t: &Self::T) -> usize {
629        t.len()
630    }
631    fn transform(mut t: Self::T, len: usize) -> Self::F {
632        t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
633        #[cfg(target_arch = "x86_64")]
634        if use_block_ntt::<M>(t.len()) {
635            unsafe { ntt_simd::transform_blocks_avx2(&mut t) };
636            return t;
637        }
638        ntt(&mut t);
639        t
640    }
641    fn inverse_transform(mut f: Self::F, len: usize) -> Self::T {
642        #[cfg(target_arch = "x86_64")]
643        if use_block_ntt::<M>(f.len()) {
644            unsafe { ntt_simd::inverse_transform_blocks_avx2(&mut f) };
645            f.truncate(len);
646            return f;
647        }
648        intt(&mut f);
649        f.truncate(len);
650        f
651    }
652    fn multiply(f: &mut Self::F, g: &Self::F) {
653        assert_eq!(f.len(), g.len());
654        #[cfg(target_arch = "x86_64")]
655        if use_block_ntt::<M>(f.len()) {
656            unsafe { ntt_simd::multiply_blocks_avx2(f, g) };
657            return;
658        }
659        pointwise_multiply(f, g);
660    }
661    fn square(t: Self::T, len: usize) -> Self::T {
662        let mut f = Self::transform(t, len);
663        let g = f.clone();
664        Self::multiply(&mut f, &g);
665        Self::inverse_transform(f, len)
666    }
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    }
743}
744
745type MVec<M> = Vec<MInt<M>>;
746
747fn convert_crt_input<M, N1, N2, N3>(t: MVec<M>, capacity: usize) -> (MVec<N1>, MVec<N2>, MVec<N3>)
748where
749    M: MIntConvert<u32>,
750    N1: Montgomery32NttModulus,
751    N2: Montgomery32NttModulus,
752    N3: Montgomery32NttModulus,
753{
754    let mut f = (
755        MVec::<N1>::with_capacity(capacity),
756        MVec::<N2>::with_capacity(capacity),
757        MVec::<N3>::with_capacity(capacity),
758    );
759    advise_huge_pages(&mut f.0);
760    advise_huge_pages(&mut f.1);
761    advise_huge_pages(&mut f.2);
762    for t in t {
763        let t: u32 = t.into();
764        f.0.push(t.into());
765        f.1.push(t.into());
766        f.2.push(t.into());
767    }
768    f
769}
770
771fn reconstruct_mint_crt<M, N1, N2, N3>(f: (MVec<N1>, MVec<N2>, MVec<N3>)) -> MVec<M>
772where
773    M: MIntConvert + MIntConvert<u32>,
774    N1: Montgomery32NttModulus,
775    N2: Montgomery32NttModulus,
776    N3: Montgomery32NttModulus,
777{
778    let t1 = MInt::<N2>::new(N1::get_mod()).inv();
779    let m1_3 = MInt::<N3>::new(N1::get_mod());
780    let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
781    let modulus = <M as MIntConvert<u32>>::mod_into() as u64;
782    let m1 = N1::get_mod() as u64;
783    let m2 = m1 * N2::get_mod() as u64 % modulus;
784    let fits_u64 = (N1::get_mod() - 1) as u128
785        + (N2::get_mod() - 1) as u128 * m1 as u128
786        + (N3::get_mod() - 1) as u128 * m2 as u128
787        <= u64::MAX as u128;
788    f.0.into_iter()
789        .zip(f.1)
790        .zip(f.2)
791        .map(|((c1, c2), c3)| {
792            let d1 = c1.inner();
793            let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
794            let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
795            let d3 = ((c3 - x) * t2).inner();
796            let value = if fits_u64 {
797                (d1 as u64 + d2 as u64 * m1 + d3 as u64 * m2) % modulus
798            } else {
799                ((d1 as u128 + d2 as u128 * m1 as u128 + d3 as u128 * m2 as u128) % modulus as u128)
800                    as u64
801            };
802            MInt::<M>::from(value as u32)
803        })
804        .collect()
805}
806
807impl<M, N1, N2, N3> ConvolveSteps for Convolve<(M, (N1, N2, N3))>
808where
809    M: MIntConvert + MIntConvert<u32>,
810    N1: Montgomery32NttModulus,
811    N2: Montgomery32NttModulus,
812    N3: Montgomery32NttModulus,
813{
814    type T = MVec<M>;
815    type F = (MVec<N1>, MVec<N2>, MVec<N3>);
816    fn length(t: &Self::T) -> usize {
817        t.len()
818    }
819    fn transform(t: Self::T, len: usize) -> Self::F {
820        let npot = len.max(1).next_power_of_two();
821        let f = convert_crt_input(t, npot);
822        (
823            Convolve::<N1>::transform(f.0, npot),
824            Convolve::<N2>::transform(f.1, npot),
825            Convolve::<N3>::transform(f.2, npot),
826        )
827    }
828    fn inverse_transform(f: Self::F, len: usize) -> Self::T {
829        reconstruct_mint_crt((
830            Convolve::<N1>::inverse_transform(f.0, len),
831            Convolve::<N2>::inverse_transform(f.1, len),
832            Convolve::<N3>::inverse_transform(f.2, len),
833        ))
834    }
835    fn multiply(f: &mut Self::F, g: &Self::F) {
836        Convolve::<N1>::multiply(&mut f.0, &g.0);
837        Convolve::<N2>::multiply(&mut f.1, &g.1);
838        Convolve::<N3>::multiply(&mut f.2, &g.2);
839    }
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    }
881}
882
883fn convolve_mint_crt<M, N1, N2, N3>(a: MVec<M>, b: MVec<M>) -> MVec<M>
884where
885    M: MIntConvert + MIntConvert<u32>,
886    N1: Montgomery32NttModulus,
887    N2: Montgomery32NttModulus,
888    N3: Montgomery32NttModulus,
889{
890    let convolve = |a: MVec<M>, b: MVec<M>| {
891        let a_len = a.len();
892        let b_len = b.len();
893        let a = convert_crt_input(a, a_len);
894        let b = convert_crt_input(b, b_len);
895        reconstruct_mint_crt((
896            Convolve::<N1>::convolve(a.0, b.0),
897            Convolve::<N2>::convolve(a.1, b.1),
898            Convolve::<N3>::convolve(a.2, b.2),
899        ))
900    };
901    let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
902    let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
903    if a.len().min(b.len()) as u128 * (modulus - 1).pow(2) < capacity {
904        return convolve(a, b);
905    }
906    let block_len = ((capacity - 1) / (modulus - 1).pow(2)) as usize;
907    if block_len == 0 {
908        return convolve_naive(&a, &b);
909    }
910    let mut result = vec![MInt::<M>::zero(); a.len() + b.len() - 1];
911    for (i, a) in a.chunks(block_len).enumerate() {
912        for (j, b) in b.chunks(block_len).enumerate() {
913            let product = convolve(a.to_vec(), b.to_vec());
914            for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
915                *value += product;
916            }
917        }
918    }
919    result
920}
921
922impl<N1, N2, N3> ConvolveSteps for Convolve<(u64, (N1, N2, N3))>
923where
924    N1: Montgomery32NttModulus,
925    N2: Montgomery32NttModulus,
926    N3: Montgomery32NttModulus,
927{
928    type T = Vec<u64>;
929    type F = ([MVec<N1>; 3], [MVec<N2>; 3], [MVec<N3>; 3]);
930
931    fn length(t: &Self::T) -> usize {
932        t.len()
933    }
934
935    fn transform(t: Self::T, len: usize) -> Self::F {
936        let npot = len.max(1).next_power_of_two();
937        assert!(npot <= 1usize << N1::RANK.min(N2::RANK).min(N3::RANK));
938        // The 22-bit fallback needs room for three limb products per coefficient.
939        assert!(
940            3 * npot as u128 * ((1u128 << 22) - 1).pow(2)
941                < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
942        );
943        let bits = if 2 * npot as u128 * (u32::MAX as u128).pow(2)
944            < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
945        {
946            32
947        } else {
948            22
949        };
950        let parts = if bits == 32 && t.iter().all(|&value| value <= u32::MAX as u64) {
951            1
952        } else {
953            64usize.div_ceil(bits)
954        };
955        fn split<M: Montgomery32NttModulus>(
956            t: &[u64],
957            len: usize,
958            bits: usize,
959            parts: usize,
960        ) -> [MVec<M>; 3] {
961            std::array::from_fn(|part| {
962                if part >= parts {
963                    return Vec::new();
964                }
965                Convolve::<M>::transform(
966                    t.iter()
967                        .map(|&t| MInt::from((t >> (part * bits)) & ((1u64 << bits) - 1)))
968                        .collect(),
969                    len,
970                )
971            })
972        }
973        (
974            split(&t, npot, bits, parts),
975            split(&t, npot, bits, parts),
976            split(&t, npot, bits, parts),
977        )
978    }
979
980    fn inverse_transform(f: Self::F, len: usize) -> Self::T {
981        let bits = if f.0[2].is_empty() { 32 } else { 22 };
982        let t1 = MInt::<N2>::new(N1::get_mod()).inv();
983        let m1 = N1::get_mod() as u64;
984        let m1_3 = MInt::<N3>::new(N1::get_mod());
985        let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
986        let m2 = m1 * N2::get_mod() as u64;
987        let mut result = vec![0u64; len.min(f.0[0].len())];
988        for (part, ((f1, f2), f3)) in f.0.into_iter().zip(f.1).zip(f.2).enumerate() {
989            if f1.is_empty() {
990                continue;
991            }
992            for (value, ((c1, c2), c3)) in result.iter_mut().zip(
993                Convolve::<N1>::inverse_transform(f1, len)
994                    .into_iter()
995                    .zip(Convolve::<N2>::inverse_transform(f2, len))
996                    .zip(Convolve::<N3>::inverse_transform(f3, len)),
997            ) {
998                let d1 = c1.inner();
999                let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
1000                let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
1001                let d3 = ((c3 - x) * t2).inner();
1002                let limb = (d1 as u64)
1003                    .wrapping_add((d2 as u64).wrapping_mul(m1))
1004                    .wrapping_add((d3 as u64).wrapping_mul(m2));
1005                *value = value.wrapping_add(limb << (part * bits));
1006            }
1007        }
1008        result
1009    }
1010
1011    fn multiply(f: &mut Self::F, g: &Self::F) {
1012        fn multiply<M: Montgomery32NttModulus>(f: &mut [MVec<M>; 3], g: &[MVec<M>; 3]) {
1013            assert_eq!(f[0].len(), g[0].len());
1014            if f[1].is_empty() || g[1].is_empty() {
1015                if f[1].is_empty() && !g[1].is_empty() {
1016                    f[1] = f[0].clone();
1017                    Convolve::<M>::multiply(&mut f[1], &g[1]);
1018                } else if !f[1].is_empty() {
1019                    Convolve::<M>::multiply(&mut f[1], &g[0]);
1020                }
1021                Convolve::<M>::multiply(&mut f[0], &g[0]);
1022                return;
1023            }
1024            #[cfg(target_arch = "x86_64")]
1025            if use_block_ntt::<M>(f[0].len()) {
1026                for part in (1..if f[2].is_empty() { 2 } else { 3 }).rev() {
1027                    let mut sum = f[0].clone();
1028                    Convolve::<M>::multiply(&mut sum, &g[part]);
1029                    for left in 1..=part {
1030                        let mut product = f[left].clone();
1031                        Convolve::<M>::multiply(&mut product, &g[part - left]);
1032                        for (value, product) in sum.iter_mut().zip(product) {
1033                            // Block products contain lazy Montgomery residues.
1034                            *value = MInt::new(value.inner() + product.inner());
1035                        }
1036                    }
1037                    f[part] = sum;
1038                }
1039                Convolve::<M>::multiply(&mut f[0], &g[0]);
1040                return;
1041            }
1042            if f[2].is_empty() {
1043                for i in 0..f[0].len() {
1044                    f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1045                    f[0][i] *= g[0][i];
1046                }
1047                return;
1048            }
1049            for i in 0..f[0].len() {
1050                f[2][i] = f[0][i] * g[2][i] + f[1][i] * g[1][i] + f[2][i] * g[0][i];
1051                f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1052                f[0][i] *= g[0][i];
1053            }
1054        }
1055        multiply(&mut f.0, &g.0);
1056        multiply(&mut f.1, &g.1);
1057        multiply(&mut f.2, &g.2);
1058    }
1059
1060    fn square(t: Self::T, len: usize) -> Self::T {
1061        let mut f = Self::transform(t, len);
1062        let g = f.clone();
1063        Self::multiply(&mut f, &g);
1064        Self::inverse_transform(f, len)
1065    }
1066
1067    fn convolve(a: Self::T, b: Self::T) -> Self::T {
1068        let max_len = Self::length(&a).max(Self::length(&b));
1069        let min_len = Self::length(&a).min(Self::length(&b));
1070        let (balanced, short) = crate::avx_helper!(@dispatch_avx2_fma (300, 64), (1536, 512));
1071        if max_len <= balanced || min_len <= short {
1072            let a_wrapping: &[Wrapping<u64>] =
1073                unsafe { std::slice::from_raw_parts(a.as_ptr().cast(), a.len()) };
1074            let b_wrapping: &[Wrapping<u64>] =
1075                unsafe { std::slice::from_raw_parts(b.as_ptr().cast(), b.len()) };
1076            let mut c = std::mem::ManuallyDrop::new(if max_len <= 300 || min_len > 60 {
1077                convolve_karatsuba(a_wrapping, b_wrapping)
1078            } else {
1079                convolve_naive(a_wrapping, b_wrapping)
1080            });
1081            return unsafe { Vec::from_raw_parts(c.as_mut_ptr().cast(), c.len(), c.capacity()) };
1082        }
1083        let len = (Self::length(&a) + Self::length(&b)).saturating_sub(1);
1084        let block_len = if min_len >= 1 << 20 {
1085            1 << 20
1086        } else {
1087            (min_len.next_power_of_two() * 8).min(1 << 21) - min_len + 1
1088        };
1089        if max_len <= block_len {
1090            return convolve_u64_fft(a, b);
1091        }
1092        let mut result = vec![0u64; len];
1093        for (i, a) in a.chunks(block_len).enumerate() {
1094            for (j, b) in b.chunks(block_len).enumerate() {
1095                if a.len().min(b.len()) <= 60 {
1096                    for (x, &a) in a.iter().enumerate() {
1097                        for (y, &b) in b.iter().enumerate() {
1098                            let value = &mut result[(i + j) * block_len + x + y];
1099                            *value = value.wrapping_add(a.wrapping_mul(b));
1100                        }
1101                    }
1102                    continue;
1103                }
1104                let product = convolve_u64_fft(a.to_vec(), b.to_vec());
1105                for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
1106                    *value = value.wrapping_add(product);
1107                }
1108            }
1109        }
1110        result
1111    }
1112}
1113
1114fn convolve_u64_fft(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1115    // Keep limb convolutions below 2^47 at the 2^21 FFT limit.
1116    crate::avx_helper!(@dispatch_avx2_fma return unsafe {
1117        convolve_u64_avx2(a, b)
1118    }, ());
1119    convolve_u64_fft_scalar(a, b)
1120}
1121
1122fn convolve_u64_fft_scalar(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1123    fn split(values: &[u64]) -> [Vec<i64>; 5] {
1124        let mut result = std::array::from_fn(|_| Vec::with_capacity(values.len()));
1125        for mut value in values.iter().copied() {
1126            for part in &mut result {
1127                let digit = ((value << 51) as i64) >> 51;
1128                part.push(digit);
1129                value = (value >> 13).wrapping_add(u64::from(digit < 0));
1130            }
1131        }
1132        result
1133    }
1134
1135    let len = a.len() + b.len() - 1;
1136    let transform = |values: &[u64]| {
1137        if values.iter().any(|&value| value > u32::MAX as u64) {
1138            return split(values).map(|part| ConvolveRealFft::transform(part, len));
1139        }
1140        let [a, b, c, _, _] = split(values);
1141        let a = ConvolveRealFft::transform(a, len);
1142        let size = a.len();
1143        [
1144            a,
1145            ConvolveRealFft::transform(b, len),
1146            ConvolveRealFft::transform(c, len),
1147            vec![Zero::zero(); size],
1148            vec![Zero::zero(); size],
1149        ]
1150    };
1151    let fa = transform(&a);
1152    drop(a);
1153    let fb = transform(&b);
1154    drop(b);
1155    let values: [Vec<i64>; 5] = std::array::from_fn(|part| {
1156        let mut sum = fa[0].clone();
1157        ConvolveRealFft::multiply(&mut sum, &fb[part]);
1158        for left in 1..=part {
1159            let mut product = fa[left].clone();
1160            ConvolveRealFft::multiply(&mut product, &fb[part - left]);
1161            for (sum, product) in sum.iter_mut().zip(product) {
1162                *sum += product;
1163            }
1164        }
1165        ConvolveRealFft::inverse_transform(sum, len)
1166    });
1167    (0..len)
1168        .map(|i| {
1169            (values[0][i] as u64)
1170                .wrapping_add((values[1][i] as u64) << 13)
1171                .wrapping_add((values[2][i] as u64) << 26)
1172                .wrapping_add((values[3][i] as u64) << 39)
1173                .wrapping_add((values[4][i] as u64) << 52)
1174        })
1175        .collect()
1176}
1177
1178pub trait NttReuse: ConvolveSteps {
1179    const MULTIPLE: bool = true;
1180
1181    /// Transforms coefficients into the usual NTT frequency order.
1182    fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1183        Self::transform(t, len)
1184    }
1185
1186    /// Inverts a value produced by `transform_ntt`.
1187    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1188        Self::inverse_transform(f, len)
1189    }
1190
1191    /// Extends a value produced by `transform_ntt` to twice its length.
1192    /// If `monic`, the input represents a monic degree-`n` polynomial modulo
1193    /// `x^n - 1`, where `n` is the transform length.
1194    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F;
1195
1196    /// Extracts the even coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1197    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1198
1199    /// Extracts the odd coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1200    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1201
1202    /// Multiplies a usual NTT transform by the corresponding prefix of another one.
1203    fn multiply_prefix(f: &mut Self::F, g: &Self::F);
1204
1205    /// Adds the pointwise product of two usual NTT transforms to `sum`.
1206    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F);
1207
1208    /// Maximum number of products that can be summed before reconstruction.
1209    /// Both factors must transform canonical coefficients at the supplied transform's length,
1210    /// and each cyclic product must itself be reconstructible.
1211    fn max_product_sum_count(_f: &Self::F) -> usize {
1212        if Self::MULTIPLE { 1 } else { usize::MAX }
1213    }
1214
1215    fn power_projection_step(
1216        p_flat: Self::T,
1217        q_flat: Self::T,
1218        n: usize,
1219        py: usize,
1220        qy: usize,
1221    ) -> (Self::T, Self::T) {
1222        let base = n * 2;
1223        let len_p = base * py;
1224        let len_q = base * qy;
1225        let len = (len_p + len_q - 1).max(len_q + len_q - 1);
1226        let half = len.max(1).next_power_of_two() / 2;
1227
1228        let p_fft = Self::transform_ntt(p_flat, len);
1229        let q_fft = Self::transform_ntt(q_flat, len);
1230        let pr_fft = Self::odd_mul_normal_neg(&p_fft, &q_fft);
1231        let qr_fft = Self::even_mul_normal_neg(&q_fft, &q_fft);
1232        (
1233            Self::inverse_transform_ntt(pr_fft, half),
1234            Self::inverse_transform_ntt(qr_fft, half),
1235        )
1236    }
1237}
1238
1239thread_local!(
1240    static BIT_REVERSE: UnsafeCell<Vec<Vec<usize>>> = const { UnsafeCell::new(vec![]) };
1241);
1242
1243impl<M> NttReuse for Convolve<M>
1244where
1245    M: Montgomery32NttModulus,
1246{
1247    const MULTIPLE: bool = false;
1248
1249    fn transform_ntt(mut t: Self::T, len: usize) -> Self::F {
1250        t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
1251        ntt(&mut t);
1252        t
1253    }
1254
1255    fn inverse_transform_ntt(mut f: Self::F, len: usize) -> Self::T {
1256        intt(&mut f);
1257        f.truncate(len);
1258        f
1259    }
1260
1261    fn ntt_doubling(mut f: Self::F, monic: bool) -> Self::F {
1262        let n = f.len();
1263        let k = n.trailing_zeros() as usize;
1264        let mut a = Self::inverse_transform_ntt(f.clone(), n);
1265        if monic {
1266            a[0] -= MInt::<M>::from(2);
1267        }
1268        let zeta = MInt::<M>::new_unchecked(M::INFO.root[k + 1]);
1269        let zeta2 = zeta * zeta;
1270        let mut rot = [MInt::one(), zeta, zeta2, zeta2 * zeta];
1271        let step = zeta2 * zeta2;
1272        for a in a.chunks_mut(4) {
1273            for (a, rot) in a.iter_mut().zip(&mut rot) {
1274                *a *= *rot;
1275                *rot *= step;
1276            }
1277        }
1278        f.extend(Self::transform_ntt(a, n));
1279        f
1280    }
1281
1282    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1283        assert_eq!(f.len(), g.len());
1284        assert!(f.len().is_power_of_two());
1285        assert!(f.len() >= 2);
1286        if std::ptr::eq(f, g) {
1287            return f.as_chunks::<2>().0.iter().map(|a| a[0] * a[1]).collect();
1288        }
1289        let inv2 = MInt::<M>::from(2).inv();
1290        let n = f.len() / 2;
1291        (0..n)
1292            .map(|i| (f[i << 1] * g[i << 1 | 1] + f[i << 1 | 1] * g[i << 1]) * inv2)
1293            .collect()
1294    }
1295
1296    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1297        assert_eq!(f.len(), g.len());
1298        assert!(f.len().is_power_of_two());
1299        assert!(f.len() >= 2);
1300        let mut inv2 = MInt::<M>::from(2).inv();
1301        let n = f.len() / 2;
1302        let k = f.len().trailing_zeros() as usize;
1303        let mut h = vec![MInt::<M>::zero(); n];
1304        let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1305        BIT_REVERSE.with(|br| {
1306            let br = unsafe { &mut *br.get() };
1307            if br.len() < k {
1308                br.resize_with(k, Default::default);
1309            }
1310            let k = k - 1;
1311            if br[k].is_empty() {
1312                let mut v = vec![0; 1 << k];
1313                for i in 0..1 << k {
1314                    v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1315                }
1316                br[k] = v;
1317            }
1318            for &i in &br[k] {
1319                h[i] = (f[i << 1] * g[i << 1 | 1] - f[i << 1 | 1] * g[i << 1]) * inv2;
1320                inv2 *= w;
1321            }
1322        });
1323        h
1324    }
1325
1326    fn multiply_prefix(f: &mut Self::F, g: &Self::F) {
1327        pointwise_multiply(f, g);
1328    }
1329
1330    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F) {
1331        assert!(sum.len() == f.len() && sum.len() == g.len());
1332        pointwise_multiply_add(sum, f, g);
1333    }
1334
1335    fn power_projection_step(
1336        p_flat: Vec<MInt<M>>,
1337        q_flat: Vec<MInt<M>>,
1338        n: usize,
1339        py: usize,
1340        qy: usize,
1341    ) -> (Vec<MInt<M>>, Vec<MInt<M>>) {
1342        let high_degree = (qy - 1) * 2;
1343        let rows = (py + qy - 1).max(high_degree).next_power_of_two();
1344        let cols = n * 2;
1345        let size = rows * cols;
1346        let mut p = p_flat;
1347        p.resize_with(size, MInt::<M>::zero);
1348        ntt_rows(&mut p, cols);
1349        ntt_batch(&mut p, cols);
1350
1351        let mut q = q_flat;
1352        q.resize_with(size, MInt::<M>::zero);
1353        ntt_rows(&mut q, cols);
1354        let q_high = (rows == high_degree).then(|| q[(qy - 1) * cols..qy * cols].to_vec());
1355        ntt_batch(&mut q, cols);
1356
1357        let half = cols / 2;
1358        let mut odd_factor = vec![MInt::<M>::zero(); half];
1359        let mut factor = MInt::<M>::from(2).inv();
1360        let k = cols.trailing_zeros() as usize;
1361        let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1362        BIT_REVERSE.with(|br| {
1363            let br = unsafe { &mut *br.get() };
1364            if br.len() < k {
1365                br.resize_with(k, Default::default);
1366            }
1367            let k = k - 1;
1368            if br[k].is_empty() {
1369                let mut v = vec![0; 1 << k];
1370                for i in 0..1 << k {
1371                    v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1372                }
1373                br[k] = v;
1374            }
1375            for &i in &br[k] {
1376                odd_factor[i] = factor;
1377                factor *= w;
1378            }
1379        });
1380
1381        let mut pr = vec![MInt::<M>::zero(); rows * half];
1382        let mut qr = vec![MInt::<M>::zero(); rows * half];
1383        for i in 0..pr.len() {
1384            pr[i] = (p[i << 1] * q[i << 1 | 1] - p[i << 1 | 1] * q[i << 1])
1385                * odd_factor[i & (half - 1)];
1386            qr[i] = q[i << 1] * q[i << 1 | 1];
1387        }
1388        intt_batch(&mut pr, half);
1389        intt_rows(&mut pr, half);
1390        intt_batch(&mut qr, half);
1391        intt_rows(&mut qr, half);
1392
1393        if let Some(q_high) = q_high {
1394            let mut q_high_even = vec![MInt::<M>::zero(); half];
1395            for i in 0..half {
1396                q_high_even[i] = q_high[i << 1] * q_high[i << 1 | 1];
1397            }
1398            intt(&mut q_high_even);
1399            for (value, high) in qr.iter_mut().zip(&q_high_even) {
1400                *value -= *high;
1401            }
1402            qr.extend_from_slice(&q_high_even);
1403        }
1404        (pr, qr)
1405    }
1406}
1407
1408impl<M, N1, N2, N3> NttReuse for Convolve<(M, (N1, N2, N3))>
1409where
1410    M: MIntConvert + MIntConvert<u32>,
1411    N1: Montgomery32NttModulus,
1412    N2: Montgomery32NttModulus,
1413    N3: Montgomery32NttModulus,
1414{
1415    fn max_product_sum_count(f: &Self::F) -> usize {
1416        let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
1417        if modulus == 1 {
1418            return usize::MAX;
1419        }
1420        let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
1421        ((capacity - 1) / ((modulus - 1) * (modulus - 1)) / f.0.len() as u128)
1422            .clamp(1, usize::MAX as u128) as usize
1423    }
1424
1425    fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1426        let npot = len.max(1).next_power_of_two();
1427        let f = convert_crt_input(t, npot);
1428        (
1429            Convolve::<N1>::transform_ntt(f.0, npot),
1430            Convolve::<N2>::transform_ntt(f.1, npot),
1431            Convolve::<N3>::transform_ntt(f.2, npot),
1432        )
1433    }
1434
1435    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1436        reconstruct_mint_crt((
1437            Convolve::<N1>::inverse_transform_ntt(f.0, len),
1438            Convolve::<N2>::inverse_transform_ntt(f.1, len),
1439            Convolve::<N3>::inverse_transform_ntt(f.2, len),
1440        ))
1441    }
1442
1443    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F {
1444        if monic {
1445            let n = f.0.len();
1446            let mut coefficients = Self::inverse_transform_ntt(f, n);
1447            coefficients[0] -= MInt::<M>::one();
1448            coefficients.push(MInt::<M>::one());
1449            Self::transform_ntt(coefficients, n * 2)
1450        } else {
1451            (
1452                Convolve::<N1>::ntt_doubling(f.0, false),
1453                Convolve::<N2>::ntt_doubling(f.1, false),
1454                Convolve::<N3>::ntt_doubling(f.2, false),
1455            )
1456        }
1457    }
crates/competitive/src/math/formal_power_series/formal_power_series_impls.rs (line 424)
398    fn exp_or_pow(&self, power: Option<T>, deg: usize) -> Self
399    where
400        C: NttReuse<T = Vec<T>>,
401        C::F: Clone,
402    {
403        if deg == 1 {
404            return Self::one();
405        }
406        let indices: Vec<_> = (0..=deg).map(T::from).collect();
407        let modulus = <T::Base as MIntConvert<usize>>::mod_into();
408        let mut inv = vec![T::zero(); deg + 1];
409        inv[1] = T::one();
410        for i in 2..=deg {
411            inv[i] = -T::from(modulus / i) * &inv[modulus % i];
412        }
413        let block = deg.next_power_of_two() / 16;
414        let logarithm = if let Some(rhs) = &power {
415            self.prefix_ref(block).log(block) * rhs
416        } else {
417            self.prefix_ref(block)
418        };
419        let (kernel, mut kernel_inverse, previous_inverse_fft) =
420            logarithm.exp_newton(block, &indices, &inv);
421        if power.is_some() {
422            kernel_inverse = (self.prefix_ref(block) * &kernel).inv(block);
423        } else {
424            let mut error_fft = C::transform_ntt(kernel.data.clone(), block);
425            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
426            let error = C::inverse_transform_ntt(error_fft, block);
427            let mut error_fft =
428                C::transform_ntt(error.into_iter().skip(block / 2).collect(), block);
429            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
430            let error = C::inverse_transform_ntt(error_fft, block / 2);
431            kernel_inverse
432                .data
433                .extend(error.into_iter().take(block / 2).map(Neg::neg));
434        }
435        let kernel_data = kernel.data;
436        let kernel_inverse_data = kernel_inverse.data;
437        let kernel_inverse = C::transform(kernel_inverse_data, block * 2);
438        let kernel = C::transform(kernel_data.clone(), block * 2);
439        let blocks = deg.div_ceil(block);
440        let mut derivative_ffts = Vec::with_capacity(blocks - 1);
441        let mut polynomial_ffts = Vec::with_capacity(if power.is_some() { blocks - 1 } else { 0 });
442        for q in 1..blocks {
443            let mut values = Self::zeros(block * 2);
444            for (i, values) in values.data.chunks_mut(block).enumerate() {
445                let start = (q - i) * block;
446                for (value, x) in values
447                    .iter_mut()
448                    .zip(self.iter().skip(start).take(deg - start))
449                {
450                    *value = x.clone();
451                }
452            }
453            if power.is_some() {
454                polynomial_ffts.push(C::transform_ntt(values.data.clone(), block * 2));
455            }
456            for (i, values) in values.data.chunks_mut(block).enumerate() {
457                let start = (q - i) * block;
458                for (value, index) in values.iter_mut().zip(&indices[start..]) {
459                    *value *= index;
460                }
461            }
462            derivative_ffts.push(C::transform_ntt(values.data, block * 2));
463        }
464        let mut result = kernel_data.clone();
465        result.reserve(deg - block);
466        let mut result_ffts = Vec::with_capacity(blocks - 1);
467        for q in 1..blocks {
468            result_ffts.push(C::transform_ntt(
469                result[(q - 1) * block..q * block].to_vec(),
470                block * 2,
471            ));
472            let mut values = Self::sum_products(&derivative_ffts[..q], &result_ffts, block);
473            if let Some(rhs) = &power {
474                let product = Self::sum_products(&polynomial_ffts[..q], &result_ffts, block);
475                let factor = rhs.clone() + T::one();
476                // The power satisfies f g' = rhs f' g.
477                for (i, value) in values.iter_mut().take(deg - q * block).enumerate() {
478                    *value = value.clone() * &factor - product[i].clone() * &indices[q * block + i];
479                }
480            }
481            let mut values = C::transform(values, block * 2);
482            C::multiply(&mut values, &kernel_inverse);
483            let mut values = C::inverse_transform(values, block * 2);
484            values.truncate(block);
485            let len = block.min(deg - q * block);
486            for (i, value) in values.iter_mut().take(len).enumerate() {
487                *value *= &inv[q * block + i];
488            }
489            values[len..].fill(T::zero());
490            let mut values = C::transform(values, block * 2);
491            C::multiply(&mut values, &kernel);
492            let mut values = C::inverse_transform(values, block * 2);
493            values.truncate(len);
494            result.extend(values);
495        }
496        Self::from_vec(result)
497    }
498
499    fn exp_newton(&self, deg: usize, indices: &[T], inv: &[T]) -> (Self, Self, C::F)
500    where
501        C: NttReuse<T = Vec<T>>,
502        C::F: Clone,
503    {
504        if deg == 1 {
505            let one = Self::one();
506            return (one.clone(), one.clone(), C::transform_ntt(one.data, 1));
507        }
508        let mut f = Self::from_vec(vec![T::one(), self.coeff(1)]);
509        let mut inverse = Self::one();
510        let mut inverse_fft = C::transform_ntt(inverse.data.clone(), 2);
511        let mut m = 2;
512        while m < deg {
513            let f_fft = C::transform_ntt(f.data.clone(), 2 * m);
514
515            let previous_inverse_fft = inverse_fft;
516            let mut error_fft = previous_inverse_fft.clone();
517            C::multiply_prefix(&mut error_fft, &f_fft);
518            let mut error = C::inverse_transform_ntt(error_fft, m);
519            error[..m / 2].fill(T::zero());
520            let mut error_fft = C::transform_ntt(error, m);
521            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
522            let error = C::inverse_transform_ntt(error_fft, m);
523            inverse
524                .data
525                .extend(error.into_iter().skip(m / 2).map(Neg::neg));
526            inverse_fft = C::transform_ntt(inverse.data.clone(), 2 * m);
527
528            let mut delta = Self::from_vec(
529                self.data
530                    .iter()
531                    .take(m)
532                    .enumerate()
533                    .skip(1)
534                    .map(|(i, value)| value.clone() * &indices[i])
535                    .collect(),
536            );
537            delta.resize(m);
538            let mut delta_fft = C::transform_ntt(delta.data, m);
539            C::multiply_prefix(&mut delta_fft, &f_fft);
540            let mut delta = Self::from_vec(C::inverse_transform_ntt(delta_fft, m));
541            for i in 1..f.length() {
542                delta[i - 1] -= f[i].clone() * &indices[i];
543            }
544            delta.resize(2 * m);
545            for i in (0..m - 1).rev() {
546                delta.data[m + i] = delta.data[i].clone();
547            }
548            delta.data[..m - 1].fill(T::zero());
549            let mut delta_fft = C::transform_ntt(delta.data, 2 * m);
550            C::multiply_prefix(&mut delta_fft, &inverse_fft);
551            let mut delta = C::inverse_transform_ntt(delta_fft, 2 * m);
552            delta.pop();
553            delta.push(T::zero());
554            let target = (2 * m).min(deg);
555            for i in (1..target).rev() {
556                delta[i] = delta[i - 1].clone() * &inv[i];
557            }
558            delta[0] = T::zero();
559            delta[target..].fill(T::zero());
560            for i in m..(2 * m).min(self.length()) {
561                delta[i] += self[i].clone();
562            }
563            delta[..m].fill(T::zero());
564            let mut delta_fft = C::transform_ntt(delta, 2 * m);
565            C::multiply_prefix(&mut delta_fft, &f_fft);
566            let delta = C::inverse_transform_ntt(delta_fft, 2 * m);
567            f.data
568                .extend(delta.into_iter().skip(m).take((deg - m).min(m)));
569            m *= 2;
570        }
571        (f, inverse, inverse_fft)
572    }
573    pub fn log(&self, deg: usize) -> Self {
574        if deg == 0 {
575            return Self::zero();
576        }
577        debug_assert!(!self[0].is_zero());
578        if deg == 1 {
579            return Self::zeros(1);
580        }
581        if let Some(step) = self.sparse_stride(deg, 2) {
582            let pos: Vec<_> = self
583                .iter()
584                .take(deg)
585                .enumerate()
586                .skip(1)
587                .filter_map(|(i, x)| (!x.is_zero()).then_some(i))
588                .collect();
589            let mut derivative = Self::zeros(deg);
590            let inverse = T::one() / self[0].clone();
591            for i in (pos.first().copied().unwrap_or(deg)..deg).step_by(step) {
592                let mut value = self.coeff(i) * T::from(i);
593                for &j in &pos {
594                    if j >= i {
595                        break;
596                    }
597                    value -= self[j].clone() * &derivative[i - j];
598                }
599                derivative[i] = value * &inverse;
600            }
601            if pos.is_empty() {
602                return derivative;
603            }
604            derivative.data.remove(0);
605            return derivative.integral();
606        }
607        let n = deg - 1;
608        if n <= 64 {
609            return (self.inv(deg) * self.prefix_ref(deg).diff())
610                .prefix(n)
611                .integral();
612        }
613        let half = n.next_power_of_two() / 2;
614        let derivative = self.prefix_ref(deg).diff();
615        let inverse = C::transform(self.inv(half).data, half * 2);
616        let mut quotient = C::transform(derivative.prefix_ref(half).data, half * 2);
617        C::multiply(&mut quotient, &inverse);
618        let mut result = C::inverse_transform(quotient, half * 2);
619        result.truncate(half);
620        if n - half <= 4 {
621            let inverse = T::one() / self[0].clone();
622            for i in half..n {
623                let mut value = derivative.coeff(i);
624                for j in 1..=i.min(self.length() - 1) {
625                    value -= self[j].clone() * &result[i - j];
626                }
627                result.push(value * &inverse);
628            }
629            return Self::from_vec(result).integral();
630        }
631        let quotient = C::transform(result.clone(), half * 2);
632        let mut error = C::transform(self.prefix_ref(n).data, half * 2);
633        C::multiply(&mut error, &quotient);
634        let mut error = C::inverse_transform(error, half * 2);
635        for i in 0..n - half {
636            error[i] = derivative.coeff(half + i) - &error[half + i];
637        }
638        error.truncate(n - half);
639        let mut error = C::transform(error, half * 2);
640        C::multiply(&mut error, &inverse);
641        let error = C::inverse_transform(error, half * 2);
642        result.extend(error.into_iter().take(n - half));
643        Self::from_vec(result).integral()
644    }
645    pub fn pow(&self, rhs: usize, deg: usize) -> Self
646    where
647        C: NttReuse<T = Vec<T>>,
648        C::F: Clone,
649    {
650        if rhs == 0 {
651            return Self::from_vec(
652                once(T::one())
653                    .chain(repeat_with(T::zero))
654                    .take(deg)
655                    .collect(),
656            );
657        }
658        if rhs == 1 {
659            return self.prefix_ref(deg).resized(deg);
660        }
661        if let Some(k) = self
662            .iter()
663            .take(deg.div_ceil(rhs))
664            .position(|x| !x.is_zero())
665        {
666            let deg = deg - k * rhs;
667            let x0 = self[k].clone();
668            let mut f = (self.prefix_ref(k + deg) >> k) / &x0;
669            if let Some(step) = f.sparse_stride(deg, 12) {
670                f = f.pow_sparse1(T::from(rhs), deg, step);
671            } else if rhs <= 4 {
672                let squared = (&f * &f).prefix(deg);
673                f = match rhs {
674                    2 => squared,
675                    3 => (squared * f).prefix(deg),
676                    _ => (&squared * &squared).prefix(deg),
677                }
678                .resized(deg);
679            } else {
680                f = f.exp_or_pow(Some(T::from(rhs)), deg);
681            }
682            f *= x0.pow(rhs);
683            f <<= k * rhs;
684            f
685        } else {
686            Self::zeros(deg)
687        }
688    }
689    fn pow_sparse1(&self, rhs: T, deg: usize, step: usize) -> Self {
690        debug_assert!(!self[0].is_zero());
691        let mut pos: Vec<_> = self
692            .data
693            .iter()
694            .take(deg)
695            .enumerate()
696            .skip(1)
697            .filter(|(_, x)| !x.is_zero())
698            .map(|(i, x)| (i, x.clone(), T::from(i) * &rhs * x))
699            .collect();
700        let mut f = Self::zeros(deg);
701        f[0] = T::one();
702        if pos.is_empty() {
703            return f;
704        }
705        let mf = T::memorized_factorial(deg);
706        for (_, coefficient, _) in &mut pos {
707            *coefficient *= T::from(step);
708        }
709        for i in (pos.first().map_or(deg, |x| x.0)..deg).step_by(step) {
710            let mut tot = T::zero();
711            for (j, coefficient, weight) in &mut pos {
712                if *j > i {
713                    break;
714                }
715                tot += weight.clone() * &f[i - *j];
716                *weight -= &*coefficient;
717            }
718            f[i] = tot * T::memorized_inv(&mf, i);
719        }
720        f
721    }
722
723    fn sparse_fold(&self, sparse: impl IntoIterator<Item = (usize, T)>, deg: usize) -> T {
724        sparse
725            .into_iter()
726            .take_while(|&(i, _)| i <= deg)
727            .fold(T::zero(), |sum, (i, x)| sum + x * self.coeff(deg - i))
728    }
729
730    /// solve: $X(QF)'=\alpha P'(QF)+\beta P(Q'F)$ in $O(deg * max(nz(P), nz(Q), nz(X)))$
731    pub fn solve_sparse_differential2(
732        p: &Self,
733        q: &Self,
734        x: &Self,
735        alpha: T,
736        beta: T,
737        deg: usize,
738    ) -> Self {
739        if deg == 0 {
740            return Self::zero();
741        }
742        let collect_sparse = |p: &Self| -> Vec<(usize, T)> {
743            p.iter()
744                .enumerate()
745                .filter(|&(_, x)| !x.is_zero())
746                .map(|(i, x)| (i, x.clone()))
747                .collect()
748        };
749        assert!(q.coeff(0).is_one());
750        assert!(x.coeff(0).is_one());
751        let p = collect_sparse(p);
752        let q = collect_sparse(q);
753        let x = collect_sparse(x);
754        let diff = |p: &[(usize, T)]| -> Vec<(usize, T)> {
755            p.iter()
756                .filter(|&&(i, _)| i > 0)
757                .map(|&(i, ref x)| (i - 1, x.clone() * T::from(i)))
758                .collect()
759        };
760        let dp = diff(&p);
761        let dq = diff(&q);
762
763        let mf = T::memorized_factorial(deg);
764        let mut f = Self::zeros(deg);
765        let mut qf = Self::zeros(deg);
766        let mut dq_f = Self::zeros(deg);
767        let mut d_qf = Self::zeros(deg);
768        f[0] = T::one();
769        for i in 0..deg - 1 {
770            qf[i] = f.sparse_fold(q.iter().cloned(), i);
771            dq_f[i] = f.sparse_fold(dq.iter().cloned(), i);
772            let dp_qf_i = qf.sparse_fold(dp.iter().cloned(), i);
773            let p_dq_f_i = dq_f.sparse_fold(p.iter().cloned(), i);
774            let x_d_qf_i = d_qf.sparse_fold(
775                x.iter()
776                    .map(|&(i, ref x)| (i, x.clone() - T::from((i == 0) as usize))),
777                i,
778            );
779            d_qf[i] = alpha.clone() * dp_qf_i + beta.clone() * p_dq_f_i - x_d_qf_i;
780
781            let mut f_ip1 = d_qf[i].clone();
782            for &(j, ref q) in q.iter().take_while(|&&(j, _)| j <= i) {
783                if j > 0 {
784                    f_ip1 -= q.clone() * &f[i - (j - 1)] * T::from(i - (j - 1));
785                }
786            }
787            f[i + 1] = f_ip1 * T::memorized_inv(&mf, i + 1);
788        }
789        f
790    }
791
792    /// P^exp_p * Q^exp_q
793    pub fn mul_of_pow_sparse(&self, q: &Self, exp_p: isize, exp_q: isize, deg: usize) -> Self {
794        if deg == 0 {
795            return Self::zero();
796        }
797        if exp_p == 0 && exp_q == 0 {
798            return Self::from_vec(
799                once(T::one())
800                    .chain(repeat_with(T::zero))
801                    .take(deg)
802                    .collect(),
803            );
804        }
805        if exp_p != 0 && self.iter().all(|x| x.is_zero()) {
806            assert!(exp_p > 0);
807            return Self::zeros(deg);
808        }
809        if exp_q != 0 && q.iter().all(|x| x.is_zero()) {
810            assert!(exp_q > 0);
811            return Self::zeros(deg);
812        }
813
814        let normalize = |f: &Self, exp: isize| {
815            if exp == 0 {
816                return (0usize, T::one(), Self::from_vec(vec![T::one()]));
817            }
818            let k = f.iter().position(|value| !value.is_zero()).unwrap();
819            assert!(
820                exp >= 0 || k == 0,
821                "Negative exponent with zero constant term"
822            );
823            let c = f[k].clone();
824            let f = (f.clone() >> k) / &c;
825            (k, c, f)
826        };
827        let (sp, cp, mut p) = normalize(self, exp_p);
828        let (sq, cq, mut q) = normalize(q, exp_q);
829
830        let shift = exp_p
831            .saturating_mul(sp as _)
832            .saturating_add(exp_q.saturating_mul(sq as _)) as usize;
833        if shift >= deg {
834            return Self::zeros(deg);
835        }
836        p.truncate(deg - shift);
837        q.truncate(deg - shift);
838
839        let mut f = Self::solve_sparse_differential2(
840            &p,
841            &q,
842            &p,
843            T::from(exp_p),
844            T::from(exp_q),
845            deg - shift,
846        );
847        f *= cp.signed_pow(exp_p) * cq.signed_pow(exp_q);
848        if shift > 0 {
849            f <<= shift;
850        }
851        f.prefix(deg)
852    }
853
854    /// exp(P/Q)
855    pub fn exp_of_div_sparse(&self, q: &Self, deg: usize) -> Self {
856        if deg == 0 {
857            return Self::zero();
858        }
859        let shift_q = q
860            .iter()
861            .position(|value| !value.is_zero())
862            .expect("Zero denominator");
863        let shift_p = self.iter().position(|value| !value.is_zero()).unwrap_or(!0);
864        assert!(shift_p > shift_q);
865
866        let mut p = self >> shift_q;
867        let mut q = q >> shift_q;
868        assert!(!q.coeff(0).is_zero());
869
870        let c = q[0].clone();
871        p /= c.clone();
872        q /= c;
873
874        Self::solve_sparse_differential2(&p, &q, &q, T::one(), -T::one(), deg)
875    }
876}
877
878impl<T, C> FormalPowerSeries<T, C>
879where
880    T: FormalPowerSeriesCoefficientSqrt,
881    C: ConvolveSteps<T = Vec<T>>,
882{
883    pub fn sqrt(&self, deg: usize) -> Option<Self> {
884        if self[0].is_zero() {
885            if let Some(k) = self.iter().position(|x| !x.is_zero()) {
886                if k % 2 != 0 {
887                    return None;
888                } else if deg > k / 2 {
889                    return Some((self >> k).sqrt(deg - k / 2)? << (k / 2));
890                }
891            }
892        } else {
893            let s = self[0].sqrt_coefficient()?;
894            if deg <= 1 {
895                return Some(Self::from(s).prefix(deg));
896            }
897            if let Some(step) = self.sparse_stride(deg, 4) {
898                let t = self[0].clone();
899                let mut f = self.prefix_ref(deg) / t;
900                f = f.pow_sparse1(T::one() / T::from(2usize), deg, step);
901                f *= s;
902                return Some(f);
903            }
904
905            let mut f = Self::from(s);
906            let inv2 = T::one() / (T::one() + T::one());
907            let inv2s = inv2.clone() / &f[0];
908            let extend = |f: &mut Self, end| {
909                for i in f.length()..end {
910                    let mut value = self.coeff(i);
911                    for j in 1..i {
912                        value -= f[j].clone() * &f[i - j];
913                    }
914                    f.data.push(value * &inv2s);
915                }
916            };
917            extend(&mut f, deg.min(32));
918            f.truncate(deg);
919            if f.length() == deg {
920                return Some(f);
921            }
922            let mut inverse = f.inv(f.length());
923            let mut i = f.length();
924            while i < deg {
925                if deg - i <= 4 {
926                    extend(&mut f, deg);
927                    break;
928                }
929                let len = (i * 2).min(deg);
930                let factor = C::transform(inverse.data.clone(), i * 2);
931                let error = if !C::CYCLIC || i < 128 {
932                    (self.prefix_ref(len) - &f * &f) >> i
933                } else {
934                    let square = C::square(f.data.clone(), i);
935                    // The cyclic square folds its high half into the already known low half.
936                    Self::from_vec(
937                        square
938                            .into_iter()
939                            .take(len - i)
940                            .enumerate()
941                            .map(|(j, value)| self.coeff(i + j) + self.coeff(j) - value)
942                            .collect(),
943                    )
944                };
945                let mut error_fft = C::transform(error.data, i * 2);
946                C::multiply(&mut error_fft, &factor);
947                let delta = C::inverse_transform(error_fft, i * 2);
948                f.data
949                    .extend(delta.into_iter().take(len - i).map(|x| x * &inv2));
950                if i * 2 + 4 < deg {
951                    let mut error_fft = C::transform(f.data.clone(), i * 2);
952                    C::multiply(&mut error_fft, &factor);
953                    let error = C::inverse_transform(error_fft, i * 2);
954                    let mut error_fft = C::transform(error.into_iter().skip(i).collect(), i * 2);
955                    C::multiply(&mut error_fft, &factor);
956                    let error = C::inverse_transform(error_fft, i * 2);
957                    inverse.data.extend(error.into_iter().take(i).map(Neg::neg));
958                }
959                i *= 2;
960            }
961            f.truncate(deg);
962            return Some(f);
963        }
964        Some(Self::zeros(deg))
965    }
966}
967
968impl<T, C> FormalPowerSeries<T, C>
969where
970    T: FormalPowerSeriesCoefficient,
971    C: ConvolveSteps<T = Vec<T>>,
972{
973    pub fn count_subset_sum<F>(&self, deg: usize, mut inverse: F) -> Self
974    where
975        F: FnMut(usize) -> T,
976        C: NttReuse<T = Vec<T>>,
977        C::F: Clone,
978    {
979        let n = self.length();
980        let mut f = Self::zeros(n);
981        for i in 1..n {
982            if !self[i].is_zero() {
983                for (j, d) in (0..n).step_by(i).enumerate().skip(1) {
984                    if j & 1 != 0 {
985                        f[d] += self[i].clone() * &inverse(j);
986                    } else {
987                        f[d] -= self[i].clone() * &inverse(j);
988                    }
989                }
990            }
991        }
992        f.exp(deg)
993    }
994    pub fn count_multiset_sum<F>(&self, deg: usize, mut inverse: F) -> Self
995    where
996        F: FnMut(usize) -> T,
997        C: NttReuse<T = Vec<T>>,
998        C::F: Clone,
999    {
1000        let n = self.length();
1001        let mut f = Self::zeros(n);
1002        for i in 1..n {
1003            if !self[i].is_zero() {
1004                for (j, d) in (0..n).step_by(i).enumerate().skip(1) {
1005                    f[d] += self[i].clone() * &inverse(j);
1006                }
1007            }
1008        }
1009        f.exp(deg)
1010    }
1011    /// [x^n] P(x) / Q(x)
1012    pub fn bostan_mori(mut self, mut rhs: Self, mut n: usize) -> T
1013    where
1014        C: NttReuse<T = Vec<T>>,
1015    {
1016        let mut res = T::zero();
1017        rhs.trim_tail_zeros();
1018        if self.length() >= rhs.length() {
1019            let r = &self / &rhs;
1020            if n < r.length() {
1021                res = r[n].clone();
1022            }
1023            self -= r * &rhs;
1024            self.trim_tail_zeros();
1025        }
1026        let mut k = rhs.length().next_power_of_two();
1027        let mut p = C::transform_ntt(self.data, k * 2);
1028        let mut q = C::transform_ntt(rhs.data, k * 2);
1029        while n > 0 {
1030            let t = C::even_mul_normal_neg(&q, &q);
1031            p = if n.is_multiple_of(2) {
1032                C::even_mul_normal_neg(&p, &q)
1033            } else {
1034                C::odd_mul_normal_neg(&p, &q)
1035            };
1036            q = t;
1037            n /= 2;
1038            if n != 0 {
1039                if n < k / 2 {
1040                    p = C::transform_ntt(C::inverse_transform_ntt(p, k / 2), k);
1041                    q = C::transform_ntt(C::inverse_transform_ntt(q, k / 2), k);
1042                    k /= 2;
1043                } else if C::MULTIPLE {
1044                    p = C::transform_ntt(C::inverse_transform_ntt(p, k), k * 2);
1045                    q = C::transform_ntt(C::inverse_transform_ntt(q, k), k * 2);
1046                } else {
1047                    p = C::ntt_doubling(p, false);
1048                    q = C::ntt_doubling(q, false);
1049                }
1050            }
1051        }
1052        let p = C::inverse_transform_ntt(p, k);
1053        let q = C::inverse_transform_ntt(q, k);
1054        res + p[0].clone() / q[0].clone()
1055    }
1056    /// return F(x) where [x^n] P(x) / Q(x) = [x^d-1] P(x) F(x)
1057    pub fn bostan_mori_msb(self, n: usize) -> Self {
1058        let d = self.length() - 1;
1059        if n == 0 {
1060            return (Self::one() << (d - 1)) / self[0].clone();
1061        }
1062        let q = self;
1063        let mq = q.clone().parity_inversion();
1064        let w = (q * &mq).even().bostan_mori_msb(n / 2);
1065        let mut s = Self::zeros(w.length() * 2 - (n % 2));
1066        for (i, x) in w.iter().enumerate() {
1067            s[i * 2 + (1 - n % 2)] = x.clone();
1068        }
1069        let len = 2 * d + 1;
1070        let ts = C::transform(s.prefix(len).data, len);
1071        mq.reversed().middle_product(&ts, len).prefix(d + 1)
1072    }
1073    /// x^n mod self
1074    pub fn pow_mod(self, n: usize) -> Self {
1075        let d = self.length() - 1;
1076        let q = self.reversed();
1077        let u = q.clone().bostan_mori_msb(n);
1078        let mut f = (u * q).prefix(d).reversed();
1079        f.trim_tail_zeros();
1080        f
1081    }
1082    fn middle_product(self, other: &C::F, deg: usize) -> Self {
1083        let n = self.length();
1084        let mut s = C::transform(self.reversed().data, deg);
1085        C::multiply(&mut s, other);
1086        Self::from_vec((C::inverse_transform(s, deg))[n - 1..].to_vec())
1087    }
1088    pub fn multipoint_evaluation(self, points: &[T]) -> Vec<T>
1089    where
1090        C: NttReuse<T = Vec<T>>,
1091        C::F: Clone,
1092    {
1093        let n = points.len();
1094        if n <= 32 || self.length() <= 32 {
1095            return points.iter().map(|p| self.eval(p.clone())).collect();
1096        }
1097        let size = n.next_power_of_two();
1098        let block = 16;
1099        let leaves = size / block;
1100        let mut subproduct_tree = Vec::with_capacity(leaves * 2);
1101        subproduct_tree.resize_with(leaves * 2, || None);
1102        let mut leaf_products = Vec::with_capacity(leaves);
1103        for i in 0..leaves {
1104            let mut product = vec![T::one()];
1105            for j in 0..block {
1106                let x = points.get(i * block + j).cloned().unwrap_or_else(T::zero);
1107                product.push(T::one());
1108                for k in (1..=j).rev() {
1109                    product[k] = product[k - 1].clone() - x.clone() * &product[k];
1110                }
1111                product[0] *= -x;
1112            }
1113            subproduct_tree[leaves + i] = Some(C::transform_ntt(product.clone(), block * 2));
1114            leaf_products.push(product);
1115        }
1116        for i in (1..leaves).rev() {
1117            let mut product = subproduct_tree[i * 2].as_ref().unwrap().clone();
1118            C::multiply_prefix(&mut product, subproduct_tree[i * 2 + 1].as_ref().unwrap());
1119            if i > 1 {
1120                product = C::ntt_doubling(product, true);
1121            }
1122            subproduct_tree[i] = Some(product);
1123        }
1124        let mut product = C::inverse_transform_ntt(subproduct_tree[1].take().unwrap(), size);
1125        product[0] -= T::one();
1126        product.push(T::one());
1127        let mut uptree_t = Vec::with_capacity(leaves * 2);
1128        uptree_t.resize_with(1, Zero::zero);
1129        let m = self.length();
1130        let v = Self::from_vec(product).reversed().resized(m);
1131        let s = C::transform(self.data, m * 2);
1132        uptree_t.push(v.inv(m).middle_product(&s, m * 2).resized(size));
1133        for i in 1..leaves {
1134            let degree = uptree_t[i].length();
1135            let spectrum = C::transform_ntt(std::mem::take(&mut uptree_t[i].data), degree);
1136            let left = subproduct_tree[i * 2].take().unwrap();
1137            let right = subproduct_tree[i * 2 + 1].take().unwrap();
1138            let mut child = spectrum.clone();
1139            C::multiply_prefix(&mut child, &right);
1140            let mut child = C::inverse_transform_ntt(child, degree);
1141            child.drain(..degree / 2);
1142            uptree_t.push(Self::from_vec(child));
1143            let mut child = spectrum;
1144            C::multiply_prefix(&mut child, &left);
1145            let mut child = C::inverse_transform_ntt(child, degree);
1146            child.drain(..degree / 2);
1147            uptree_t.push(Self::from_vec(child));
1148        }
1149        let mut result = Vec::with_capacity(n);
1150        for ((values, product), points) in uptree_t[leaves..]
1151            .iter()
1152            .zip(leaf_products)
1153            .zip(points.chunks(block))
1154        {
1155            let mut remainder = Self::zeros(block);
1156            for (j, value) in values.iter().enumerate() {
1157                for (r, p) in remainder.data[..=j].iter_mut().zip(&product[block - j..]) {
1158                    *r += value.clone() * p;
1159                }
1160            }
1161            result.extend(points.iter().map(|p| remainder.eval(p.clone())));
1162        }
1163        result
1164    }
Source

fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T

Inverts a value produced by transform_ntt.

Examples found in repository?
crates/competitive/src/math/formal_power_series/berlekamp_massey.rs (line 133)
131    fn apply(&self, p: &C::F, q: &C::F, length: usize) -> (Vec<T>, Vec<T>) {
132        (
133            C::inverse_transform_ntt(Self::product_sum(p, &self.a00, q, &self.a01), length),
134            C::inverse_transform_ntt(Self::product_sum(p, &self.a10, q, &self.a11), length),
135        )
136    }
137
138    fn left_multiply_step(self, quotient: &FormalPowerSeries<T, C>, length: usize) -> Self {
139        let negative_quotient = reduced_transform(&(-quotient), length);
140        let mut a10 = self.a00;
141        C::multiply_add(&mut a10, &negative_quotient, &self.a10);
142        let mut a11 = self.a01;
143        C::multiply_add(&mut a11, &negative_quotient, &self.a11);
144        let result = Self {
145            a00: self.a10,
146            a01: self.a11,
147            a10,
148            a11,
149        };
150        if C::MULTIPLE {
151            result.inverse_transform(length).transform(length)
152        } else {
153            result
154        }
155    }
156
157    fn inverse_transform(self, length: usize) -> FpsMatrix<T, C> {
158        FpsMatrix {
159            a00: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a00, length)),
160            a01: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a01, length)),
161            a10: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a10, length)),
162            a11: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a11, length)),
163        }
164    }
More examples
Hide additional examples
crates/competitive/src/math/formal_power_series/formal_power_series_impls.rs (line 387)
373    fn sum_products(f: &[C::F], g: &[C::F], len: usize) -> Vec<T>
374    where
375        C: NttReuse<T = Vec<T>>,
376        C::F: Clone,
377    {
378        let chunk = C::max_product_sum_count(&f[0]);
379        f.rchunks(chunk)
380            .zip(g.chunks(chunk))
381            .map(|(f, g)| {
382                let mut sum = f[f.len() - 1].clone();
383                C::multiply_prefix(&mut sum, &g[0]);
384                for (f, g) in f.iter().rev().skip(1).zip(&g[1..]) {
385                    C::multiply_add(&mut sum, f, g);
386                }
387                C::inverse_transform_ntt(sum, len)
388            })
389            .reduce(|mut sum, part| {
390                for (sum, value) in sum.iter_mut().zip(part) {
391                    *sum += value;
392                }
393                sum
394            })
395            .unwrap()
396    }
397
398    fn exp_or_pow(&self, power: Option<T>, deg: usize) -> Self
399    where
400        C: NttReuse<T = Vec<T>>,
401        C::F: Clone,
402    {
403        if deg == 1 {
404            return Self::one();
405        }
406        let indices: Vec<_> = (0..=deg).map(T::from).collect();
407        let modulus = <T::Base as MIntConvert<usize>>::mod_into();
408        let mut inv = vec![T::zero(); deg + 1];
409        inv[1] = T::one();
410        for i in 2..=deg {
411            inv[i] = -T::from(modulus / i) * &inv[modulus % i];
412        }
413        let block = deg.next_power_of_two() / 16;
414        let logarithm = if let Some(rhs) = &power {
415            self.prefix_ref(block).log(block) * rhs
416        } else {
417            self.prefix_ref(block)
418        };
419        let (kernel, mut kernel_inverse, previous_inverse_fft) =
420            logarithm.exp_newton(block, &indices, &inv);
421        if power.is_some() {
422            kernel_inverse = (self.prefix_ref(block) * &kernel).inv(block);
423        } else {
424            let mut error_fft = C::transform_ntt(kernel.data.clone(), block);
425            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
426            let error = C::inverse_transform_ntt(error_fft, block);
427            let mut error_fft =
428                C::transform_ntt(error.into_iter().skip(block / 2).collect(), block);
429            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
430            let error = C::inverse_transform_ntt(error_fft, block / 2);
431            kernel_inverse
432                .data
433                .extend(error.into_iter().take(block / 2).map(Neg::neg));
434        }
435        let kernel_data = kernel.data;
436        let kernel_inverse_data = kernel_inverse.data;
437        let kernel_inverse = C::transform(kernel_inverse_data, block * 2);
438        let kernel = C::transform(kernel_data.clone(), block * 2);
439        let blocks = deg.div_ceil(block);
440        let mut derivative_ffts = Vec::with_capacity(blocks - 1);
441        let mut polynomial_ffts = Vec::with_capacity(if power.is_some() { blocks - 1 } else { 0 });
442        for q in 1..blocks {
443            let mut values = Self::zeros(block * 2);
444            for (i, values) in values.data.chunks_mut(block).enumerate() {
445                let start = (q - i) * block;
446                for (value, x) in values
447                    .iter_mut()
448                    .zip(self.iter().skip(start).take(deg - start))
449                {
450                    *value = x.clone();
451                }
452            }
453            if power.is_some() {
454                polynomial_ffts.push(C::transform_ntt(values.data.clone(), block * 2));
455            }
456            for (i, values) in values.data.chunks_mut(block).enumerate() {
457                let start = (q - i) * block;
458                for (value, index) in values.iter_mut().zip(&indices[start..]) {
459                    *value *= index;
460                }
461            }
462            derivative_ffts.push(C::transform_ntt(values.data, block * 2));
463        }
464        let mut result = kernel_data.clone();
465        result.reserve(deg - block);
466        let mut result_ffts = Vec::with_capacity(blocks - 1);
467        for q in 1..blocks {
468            result_ffts.push(C::transform_ntt(
469                result[(q - 1) * block..q * block].to_vec(),
470                block * 2,
471            ));
472            let mut values = Self::sum_products(&derivative_ffts[..q], &result_ffts, block);
473            if let Some(rhs) = &power {
474                let product = Self::sum_products(&polynomial_ffts[..q], &result_ffts, block);
475                let factor = rhs.clone() + T::one();
476                // The power satisfies f g' = rhs f' g.
477                for (i, value) in values.iter_mut().take(deg - q * block).enumerate() {
478                    *value = value.clone() * &factor - product[i].clone() * &indices[q * block + i];
479                }
480            }
481            let mut values = C::transform(values, block * 2);
482            C::multiply(&mut values, &kernel_inverse);
483            let mut values = C::inverse_transform(values, block * 2);
484            values.truncate(block);
485            let len = block.min(deg - q * block);
486            for (i, value) in values.iter_mut().take(len).enumerate() {
487                *value *= &inv[q * block + i];
488            }
489            values[len..].fill(T::zero());
490            let mut values = C::transform(values, block * 2);
491            C::multiply(&mut values, &kernel);
492            let mut values = C::inverse_transform(values, block * 2);
493            values.truncate(len);
494            result.extend(values);
495        }
496        Self::from_vec(result)
497    }
498
499    fn exp_newton(&self, deg: usize, indices: &[T], inv: &[T]) -> (Self, Self, C::F)
500    where
501        C: NttReuse<T = Vec<T>>,
502        C::F: Clone,
503    {
504        if deg == 1 {
505            let one = Self::one();
506            return (one.clone(), one.clone(), C::transform_ntt(one.data, 1));
507        }
508        let mut f = Self::from_vec(vec![T::one(), self.coeff(1)]);
509        let mut inverse = Self::one();
510        let mut inverse_fft = C::transform_ntt(inverse.data.clone(), 2);
511        let mut m = 2;
512        while m < deg {
513            let f_fft = C::transform_ntt(f.data.clone(), 2 * m);
514
515            let previous_inverse_fft = inverse_fft;
516            let mut error_fft = previous_inverse_fft.clone();
517            C::multiply_prefix(&mut error_fft, &f_fft);
518            let mut error = C::inverse_transform_ntt(error_fft, m);
519            error[..m / 2].fill(T::zero());
520            let mut error_fft = C::transform_ntt(error, m);
521            C::multiply_prefix(&mut error_fft, &previous_inverse_fft);
522            let error = C::inverse_transform_ntt(error_fft, m);
523            inverse
524                .data
525                .extend(error.into_iter().skip(m / 2).map(Neg::neg));
526            inverse_fft = C::transform_ntt(inverse.data.clone(), 2 * m);
527
528            let mut delta = Self::from_vec(
529                self.data
530                    .iter()
531                    .take(m)
532                    .enumerate()
533                    .skip(1)
534                    .map(|(i, value)| value.clone() * &indices[i])
535                    .collect(),
536            );
537            delta.resize(m);
538            let mut delta_fft = C::transform_ntt(delta.data, m);
539            C::multiply_prefix(&mut delta_fft, &f_fft);
540            let mut delta = Self::from_vec(C::inverse_transform_ntt(delta_fft, m));
541            for i in 1..f.length() {
542                delta[i - 1] -= f[i].clone() * &indices[i];
543            }
544            delta.resize(2 * m);
545            for i in (0..m - 1).rev() {
546                delta.data[m + i] = delta.data[i].clone();
547            }
548            delta.data[..m - 1].fill(T::zero());
549            let mut delta_fft = C::transform_ntt(delta.data, 2 * m);
550            C::multiply_prefix(&mut delta_fft, &inverse_fft);
551            let mut delta = C::inverse_transform_ntt(delta_fft, 2 * m);
552            delta.pop();
553            delta.push(T::zero());
554            let target = (2 * m).min(deg);
555            for i in (1..target).rev() {
556                delta[i] = delta[i - 1].clone() * &inv[i];
557            }
558            delta[0] = T::zero();
559            delta[target..].fill(T::zero());
560            for i in m..(2 * m).min(self.length()) {
561                delta[i] += self[i].clone();
562            }
563            delta[..m].fill(T::zero());
564            let mut delta_fft = C::transform_ntt(delta, 2 * m);
565            C::multiply_prefix(&mut delta_fft, &f_fft);
566            let delta = C::inverse_transform_ntt(delta_fft, 2 * m);
567            f.data
568                .extend(delta.into_iter().skip(m).take((deg - m).min(m)));
569            m *= 2;
570        }
571        (f, inverse, inverse_fft)
572    }
573    pub fn log(&self, deg: usize) -> Self {
574        if deg == 0 {
575            return Self::zero();
576        }
577        debug_assert!(!self[0].is_zero());
578        if deg == 1 {
579            return Self::zeros(1);
580        }
581        if let Some(step) = self.sparse_stride(deg, 2) {
582            let pos: Vec<_> = self
583                .iter()
584                .take(deg)
585                .enumerate()
586                .skip(1)
587                .filter_map(|(i, x)| (!x.is_zero()).then_some(i))
588                .collect();
589            let mut derivative = Self::zeros(deg);
590            let inverse = T::one() / self[0].clone();
591            for i in (pos.first().copied().unwrap_or(deg)..deg).step_by(step) {
592                let mut value = self.coeff(i) * T::from(i);
593                for &j in &pos {
594                    if j >= i {
595                        break;
596                    }
597                    value -= self[j].clone() * &derivative[i - j];
598                }
599                derivative[i] = value * &inverse;
600            }
601            if pos.is_empty() {
602                return derivative;
603            }
604            derivative.data.remove(0);
605            return derivative.integral();
606        }
607        let n = deg - 1;
608        if n <= 64 {
609            return (self.inv(deg) * self.prefix_ref(deg).diff())
610                .prefix(n)
611                .integral();
612        }
613        let half = n.next_power_of_two() / 2;
614        let derivative = self.prefix_ref(deg).diff();
615        let inverse = C::transform(self.inv(half).data, half * 2);
616        let mut quotient = C::transform(derivative.prefix_ref(half).data, half * 2);
617        C::multiply(&mut quotient, &inverse);
618        let mut result = C::inverse_transform(quotient, half * 2);
619        result.truncate(half);
620        if n - half <= 4 {
621            let inverse = T::one() / self[0].clone();
622            for i in half..n {
623                let mut value = derivative.coeff(i);
624                for j in 1..=i.min(self.length() - 1) {
625                    value -= self[j].clone() * &result[i - j];
626                }
627                result.push(value * &inverse);
628            }
629            return Self::from_vec(result).integral();
630        }
631        let quotient = C::transform(result.clone(), half * 2);
632        let mut error = C::transform(self.prefix_ref(n).data, half * 2);
633        C::multiply(&mut error, &quotient);
634        let mut error = C::inverse_transform(error, half * 2);
635        for i in 0..n - half {
636            error[i] = derivative.coeff(half + i) - &error[half + i];
637        }
638        error.truncate(n - half);
639        let mut error = C::transform(error, half * 2);
640        C::multiply(&mut error, &inverse);
641        let error = C::inverse_transform(error, half * 2);
642        result.extend(error.into_iter().take(n - half));
643        Self::from_vec(result).integral()
644    }
645    pub fn pow(&self, rhs: usize, deg: usize) -> Self
646    where
647        C: NttReuse<T = Vec<T>>,
648        C::F: Clone,
649    {
650        if rhs == 0 {
651            return Self::from_vec(
652                once(T::one())
653                    .chain(repeat_with(T::zero))
654                    .take(deg)
655                    .collect(),
656            );
657        }
658        if rhs == 1 {
659            return self.prefix_ref(deg).resized(deg);
660        }
661        if let Some(k) = self
662            .iter()
663            .take(deg.div_ceil(rhs))
664            .position(|x| !x.is_zero())
665        {
666            let deg = deg - k * rhs;
667            let x0 = self[k].clone();
668            let mut f = (self.prefix_ref(k + deg) >> k) / &x0;
669            if let Some(step) = f.sparse_stride(deg, 12) {
670                f = f.pow_sparse1(T::from(rhs), deg, step);
671            } else if rhs <= 4 {
672                let squared = (&f * &f).prefix(deg);
673                f = match rhs {
674                    2 => squared,
675                    3 => (squared * f).prefix(deg),
676                    _ => (&squared * &squared).prefix(deg),
677                }
678                .resized(deg);
679            } else {
680                f = f.exp_or_pow(Some(T::from(rhs)), deg);
681            }
682            f *= x0.pow(rhs);
683            f <<= k * rhs;
684            f
685        } else {
686            Self::zeros(deg)
687        }
688    }
689    fn pow_sparse1(&self, rhs: T, deg: usize, step: usize) -> Self {
690        debug_assert!(!self[0].is_zero());
691        let mut pos: Vec<_> = self
692            .data
693            .iter()
694            .take(deg)
695            .enumerate()
696            .skip(1)
697            .filter(|(_, x)| !x.is_zero())
698            .map(|(i, x)| (i, x.clone(), T::from(i) * &rhs * x))
699            .collect();
700        let mut f = Self::zeros(deg);
701        f[0] = T::one();
702        if pos.is_empty() {
703            return f;
704        }
705        let mf = T::memorized_factorial(deg);
706        for (_, coefficient, _) in &mut pos {
707            *coefficient *= T::from(step);
708        }
709        for i in (pos.first().map_or(deg, |x| x.0)..deg).step_by(step) {
710            let mut tot = T::zero();
711            for (j, coefficient, weight) in &mut pos {
712                if *j > i {
713                    break;
714                }
715                tot += weight.clone() * &f[i - *j];
716                *weight -= &*coefficient;
717            }
718            f[i] = tot * T::memorized_inv(&mf, i);
719        }
720        f
721    }
722
723    fn sparse_fold(&self, sparse: impl IntoIterator<Item = (usize, T)>, deg: usize) -> T {
724        sparse
725            .into_iter()
726            .take_while(|&(i, _)| i <= deg)
727            .fold(T::zero(), |sum, (i, x)| sum + x * self.coeff(deg - i))
728    }
729
730    /// solve: $X(QF)'=\alpha P'(QF)+\beta P(Q'F)$ in $O(deg * max(nz(P), nz(Q), nz(X)))$
731    pub fn solve_sparse_differential2(
732        p: &Self,
733        q: &Self,
734        x: &Self,
735        alpha: T,
736        beta: T,
737        deg: usize,
738    ) -> Self {
739        if deg == 0 {
740            return Self::zero();
741        }
742        let collect_sparse = |p: &Self| -> Vec<(usize, T)> {
743            p.iter()
744                .enumerate()
745                .filter(|&(_, x)| !x.is_zero())
746                .map(|(i, x)| (i, x.clone()))
747                .collect()
748        };
749        assert!(q.coeff(0).is_one());
750        assert!(x.coeff(0).is_one());
751        let p = collect_sparse(p);
752        let q = collect_sparse(q);
753        let x = collect_sparse(x);
754        let diff = |p: &[(usize, T)]| -> Vec<(usize, T)> {
755            p.iter()
756                .filter(|&&(i, _)| i > 0)
757                .map(|&(i, ref x)| (i - 1, x.clone() * T::from(i)))
758                .collect()
759        };
760        let dp = diff(&p);
761        let dq = diff(&q);
762
763        let mf = T::memorized_factorial(deg);
764        let mut f = Self::zeros(deg);
765        let mut qf = Self::zeros(deg);
766        let mut dq_f = Self::zeros(deg);
767        let mut d_qf = Self::zeros(deg);
768        f[0] = T::one();
769        for i in 0..deg - 1 {
770            qf[i] = f.sparse_fold(q.iter().cloned(), i);
771            dq_f[i] = f.sparse_fold(dq.iter().cloned(), i);
772            let dp_qf_i = qf.sparse_fold(dp.iter().cloned(), i);
773            let p_dq_f_i = dq_f.sparse_fold(p.iter().cloned(), i);
774            let x_d_qf_i = d_qf.sparse_fold(
775                x.iter()
776                    .map(|&(i, ref x)| (i, x.clone() - T::from((i == 0) as usize))),
777                i,
778            );
779            d_qf[i] = alpha.clone() * dp_qf_i + beta.clone() * p_dq_f_i - x_d_qf_i;
780
781            let mut f_ip1 = d_qf[i].clone();
782            for &(j, ref q) in q.iter().take_while(|&&(j, _)| j <= i) {
783                if j > 0 {
784                    f_ip1 -= q.clone() * &f[i - (j - 1)] * T::from(i - (j - 1));
785                }
786            }
787            f[i + 1] = f_ip1 * T::memorized_inv(&mf, i + 1);
788        }
789        f
790    }
791
792    /// P^exp_p * Q^exp_q
793    pub fn mul_of_pow_sparse(&self, q: &Self, exp_p: isize, exp_q: isize, deg: usize) -> Self {
794        if deg == 0 {
795            return Self::zero();
796        }
797        if exp_p == 0 && exp_q == 0 {
798            return Self::from_vec(
799                once(T::one())
800                    .chain(repeat_with(T::zero))
801                    .take(deg)
802                    .collect(),
803            );
804        }
805        if exp_p != 0 && self.iter().all(|x| x.is_zero()) {
806            assert!(exp_p > 0);
807            return Self::zeros(deg);
808        }
809        if exp_q != 0 && q.iter().all(|x| x.is_zero()) {
810            assert!(exp_q > 0);
811            return Self::zeros(deg);
812        }
813
814        let normalize = |f: &Self, exp: isize| {
815            if exp == 0 {
816                return (0usize, T::one(), Self::from_vec(vec![T::one()]));
817            }
818            let k = f.iter().position(|value| !value.is_zero()).unwrap();
819            assert!(
820                exp >= 0 || k == 0,
821                "Negative exponent with zero constant term"
822            );
823            let c = f[k].clone();
824            let f = (f.clone() >> k) / &c;
825            (k, c, f)
826        };
827        let (sp, cp, mut p) = normalize(self, exp_p);
828        let (sq, cq, mut q) = normalize(q, exp_q);
829
830        let shift = exp_p
831            .saturating_mul(sp as _)
832            .saturating_add(exp_q.saturating_mul(sq as _)) as usize;
833        if shift >= deg {
834            return Self::zeros(deg);
835        }
836        p.truncate(deg - shift);
837        q.truncate(deg - shift);
838
839        let mut f = Self::solve_sparse_differential2(
840            &p,
841            &q,
842            &p,
843            T::from(exp_p),
844            T::from(exp_q),
845            deg - shift,
846        );
847        f *= cp.signed_pow(exp_p) * cq.signed_pow(exp_q);
848        if shift > 0 {
849            f <<= shift;
850        }
851        f.prefix(deg)
852    }
853
854    /// exp(P/Q)
855    pub fn exp_of_div_sparse(&self, q: &Self, deg: usize) -> Self {
856        if deg == 0 {
857            return Self::zero();
858        }
859        let shift_q = q
860            .iter()
861            .position(|value| !value.is_zero())
862            .expect("Zero denominator");
863        let shift_p = self.iter().position(|value| !value.is_zero()).unwrap_or(!0);
864        assert!(shift_p > shift_q);
865
866        let mut p = self >> shift_q;
867        let mut q = q >> shift_q;
868        assert!(!q.coeff(0).is_zero());
869
870        let c = q[0].clone();
871        p /= c.clone();
872        q /= c;
873
874        Self::solve_sparse_differential2(&p, &q, &q, T::one(), -T::one(), deg)
875    }
876}
877
878impl<T, C> FormalPowerSeries<T, C>
879where
880    T: FormalPowerSeriesCoefficientSqrt,
881    C: ConvolveSteps<T = Vec<T>>,
882{
883    pub fn sqrt(&self, deg: usize) -> Option<Self> {
884        if self[0].is_zero() {
885            if let Some(k) = self.iter().position(|x| !x.is_zero()) {
886                if k % 2 != 0 {
887                    return None;
888                } else if deg > k / 2 {
889                    return Some((self >> k).sqrt(deg - k / 2)? << (k / 2));
890                }
891            }
892        } else {
893            let s = self[0].sqrt_coefficient()?;
894            if deg <= 1 {
895                return Some(Self::from(s).prefix(deg));
896            }
897            if let Some(step) = self.sparse_stride(deg, 4) {
898                let t = self[0].clone();
899                let mut f = self.prefix_ref(deg) / t;
900                f = f.pow_sparse1(T::one() / T::from(2usize), deg, step);
901                f *= s;
902                return Some(f);
903            }
904
905            let mut f = Self::from(s);
906            let inv2 = T::one() / (T::one() + T::one());
907            let inv2s = inv2.clone() / &f[0];
908            let extend = |f: &mut Self, end| {
909                for i in f.length()..end {
910                    let mut value = self.coeff(i);
911                    for j in 1..i {
912                        value -= f[j].clone() * &f[i - j];
913                    }
914                    f.data.push(value * &inv2s);
915                }
916            };
917            extend(&mut f, deg.min(32));
918            f.truncate(deg);
919            if f.length() == deg {
920                return Some(f);
921            }
922            let mut inverse = f.inv(f.length());
923            let mut i = f.length();
924            while i < deg {
925                if deg - i <= 4 {
926                    extend(&mut f, deg);
927                    break;
928                }
929                let len = (i * 2).min(deg);
930                let factor = C::transform(inverse.data.clone(), i * 2);
931                let error = if !C::CYCLIC || i < 128 {
932                    (self.prefix_ref(len) - &f * &f) >> i
933                } else {
934                    let square = C::square(f.data.clone(), i);
935                    // The cyclic square folds its high half into the already known low half.
936                    Self::from_vec(
937                        square
938                            .into_iter()
939                            .take(len - i)
940                            .enumerate()
941                            .map(|(j, value)| self.coeff(i + j) + self.coeff(j) - value)
942                            .collect(),
943                    )
944                };
945                let mut error_fft = C::transform(error.data, i * 2);
946                C::multiply(&mut error_fft, &factor);
947                let delta = C::inverse_transform(error_fft, i * 2);
948                f.data
949                    .extend(delta.into_iter().take(len - i).map(|x| x * &inv2));
950                if i * 2 + 4 < deg {
951                    let mut error_fft = C::transform(f.data.clone(), i * 2);
952                    C::multiply(&mut error_fft, &factor);
953                    let error = C::inverse_transform(error_fft, i * 2);
954                    let mut error_fft = C::transform(error.into_iter().skip(i).collect(), i * 2);
955                    C::multiply(&mut error_fft, &factor);
956                    let error = C::inverse_transform(error_fft, i * 2);
957                    inverse.data.extend(error.into_iter().take(i).map(Neg::neg));
958                }
959                i *= 2;
960            }
961            f.truncate(deg);
962            return Some(f);
963        }
964        Some(Self::zeros(deg))
965    }
966}
967
968impl<T, C> FormalPowerSeries<T, C>
969where
970    T: FormalPowerSeriesCoefficient,
971    C: ConvolveSteps<T = Vec<T>>,
972{
973    pub fn count_subset_sum<F>(&self, deg: usize, mut inverse: F) -> Self
974    where
975        F: FnMut(usize) -> T,
976        C: NttReuse<T = Vec<T>>,
977        C::F: Clone,
978    {
979        let n = self.length();
980        let mut f = Self::zeros(n);
981        for i in 1..n {
982            if !self[i].is_zero() {
983                for (j, d) in (0..n).step_by(i).enumerate().skip(1) {
984                    if j & 1 != 0 {
985                        f[d] += self[i].clone() * &inverse(j);
986                    } else {
987                        f[d] -= self[i].clone() * &inverse(j);
988                    }
989                }
990            }
991        }
992        f.exp(deg)
993    }
994    pub fn count_multiset_sum<F>(&self, deg: usize, mut inverse: F) -> Self
995    where
996        F: FnMut(usize) -> T,
997        C: NttReuse<T = Vec<T>>,
998        C::F: Clone,
999    {
1000        let n = self.length();
1001        let mut f = Self::zeros(n);
1002        for i in 1..n {
1003            if !self[i].is_zero() {
1004                for (j, d) in (0..n).step_by(i).enumerate().skip(1) {
1005                    f[d] += self[i].clone() * &inverse(j);
1006                }
1007            }
1008        }
1009        f.exp(deg)
1010    }
1011    /// [x^n] P(x) / Q(x)
1012    pub fn bostan_mori(mut self, mut rhs: Self, mut n: usize) -> T
1013    where
1014        C: NttReuse<T = Vec<T>>,
1015    {
1016        let mut res = T::zero();
1017        rhs.trim_tail_zeros();
1018        if self.length() >= rhs.length() {
1019            let r = &self / &rhs;
1020            if n < r.length() {
1021                res = r[n].clone();
1022            }
1023            self -= r * &rhs;
1024            self.trim_tail_zeros();
1025        }
1026        let mut k = rhs.length().next_power_of_two();
1027        let mut p = C::transform_ntt(self.data, k * 2);
1028        let mut q = C::transform_ntt(rhs.data, k * 2);
1029        while n > 0 {
1030            let t = C::even_mul_normal_neg(&q, &q);
1031            p = if n.is_multiple_of(2) {
1032                C::even_mul_normal_neg(&p, &q)
1033            } else {
1034                C::odd_mul_normal_neg(&p, &q)
1035            };
1036            q = t;
1037            n /= 2;
1038            if n != 0 {
1039                if n < k / 2 {
1040                    p = C::transform_ntt(C::inverse_transform_ntt(p, k / 2), k);
1041                    q = C::transform_ntt(C::inverse_transform_ntt(q, k / 2), k);
1042                    k /= 2;
1043                } else if C::MULTIPLE {
1044                    p = C::transform_ntt(C::inverse_transform_ntt(p, k), k * 2);
1045                    q = C::transform_ntt(C::inverse_transform_ntt(q, k), k * 2);
1046                } else {
1047                    p = C::ntt_doubling(p, false);
1048                    q = C::ntt_doubling(q, false);
1049                }
1050            }
1051        }
1052        let p = C::inverse_transform_ntt(p, k);
1053        let q = C::inverse_transform_ntt(q, k);
1054        res + p[0].clone() / q[0].clone()
1055    }
1056    /// return F(x) where [x^n] P(x) / Q(x) = [x^d-1] P(x) F(x)
1057    pub fn bostan_mori_msb(self, n: usize) -> Self {
1058        let d = self.length() - 1;
1059        if n == 0 {
1060            return (Self::one() << (d - 1)) / self[0].clone();
1061        }
1062        let q = self;
1063        let mq = q.clone().parity_inversion();
1064        let w = (q * &mq).even().bostan_mori_msb(n / 2);
1065        let mut s = Self::zeros(w.length() * 2 - (n % 2));
1066        for (i, x) in w.iter().enumerate() {
1067            s[i * 2 + (1 - n % 2)] = x.clone();
1068        }
1069        let len = 2 * d + 1;
1070        let ts = C::transform(s.prefix(len).data, len);
1071        mq.reversed().middle_product(&ts, len).prefix(d + 1)
1072    }
1073    /// x^n mod self
1074    pub fn pow_mod(self, n: usize) -> Self {
1075        let d = self.length() - 1;
1076        let q = self.reversed();
1077        let u = q.clone().bostan_mori_msb(n);
1078        let mut f = (u * q).prefix(d).reversed();
1079        f.trim_tail_zeros();
1080        f
1081    }
1082    fn middle_product(self, other: &C::F, deg: usize) -> Self {
1083        let n = self.length();
1084        let mut s = C::transform(self.reversed().data, deg);
1085        C::multiply(&mut s, other);
1086        Self::from_vec((C::inverse_transform(s, deg))[n - 1..].to_vec())
1087    }
1088    pub fn multipoint_evaluation(self, points: &[T]) -> Vec<T>
1089    where
1090        C: NttReuse<T = Vec<T>>,
1091        C::F: Clone,
1092    {
1093        let n = points.len();
1094        if n <= 32 || self.length() <= 32 {
1095            return points.iter().map(|p| self.eval(p.clone())).collect();
1096        }
1097        let size = n.next_power_of_two();
1098        let block = 16;
1099        let leaves = size / block;
1100        let mut subproduct_tree = Vec::with_capacity(leaves * 2);
1101        subproduct_tree.resize_with(leaves * 2, || None);
1102        let mut leaf_products = Vec::with_capacity(leaves);
1103        for i in 0..leaves {
1104            let mut product = vec![T::one()];
1105            for j in 0..block {
1106                let x = points.get(i * block + j).cloned().unwrap_or_else(T::zero);
1107                product.push(T::one());
1108                for k in (1..=j).rev() {
1109                    product[k] = product[k - 1].clone() - x.clone() * &product[k];
1110                }
1111                product[0] *= -x;
1112            }
1113            subproduct_tree[leaves + i] = Some(C::transform_ntt(product.clone(), block * 2));
1114            leaf_products.push(product);
1115        }
1116        for i in (1..leaves).rev() {
1117            let mut product = subproduct_tree[i * 2].as_ref().unwrap().clone();
1118            C::multiply_prefix(&mut product, subproduct_tree[i * 2 + 1].as_ref().unwrap());
1119            if i > 1 {
1120                product = C::ntt_doubling(product, true);
1121            }
1122            subproduct_tree[i] = Some(product);
1123        }
1124        let mut product = C::inverse_transform_ntt(subproduct_tree[1].take().unwrap(), size);
1125        product[0] -= T::one();
1126        product.push(T::one());
1127        let mut uptree_t = Vec::with_capacity(leaves * 2);
1128        uptree_t.resize_with(1, Zero::zero);
1129        let m = self.length();
1130        let v = Self::from_vec(product).reversed().resized(m);
1131        let s = C::transform(self.data, m * 2);
1132        uptree_t.push(v.inv(m).middle_product(&s, m * 2).resized(size));
1133        for i in 1..leaves {
1134            let degree = uptree_t[i].length();
1135            let spectrum = C::transform_ntt(std::mem::take(&mut uptree_t[i].data), degree);
1136            let left = subproduct_tree[i * 2].take().unwrap();
1137            let right = subproduct_tree[i * 2 + 1].take().unwrap();
1138            let mut child = spectrum.clone();
1139            C::multiply_prefix(&mut child, &right);
1140            let mut child = C::inverse_transform_ntt(child, degree);
1141            child.drain(..degree / 2);
1142            uptree_t.push(Self::from_vec(child));
1143            let mut child = spectrum;
1144            C::multiply_prefix(&mut child, &left);
1145            let mut child = C::inverse_transform_ntt(child, degree);
1146            child.drain(..degree / 2);
1147            uptree_t.push(Self::from_vec(child));
1148        }
1149        let mut result = Vec::with_capacity(n);
1150        for ((values, product), points) in uptree_t[leaves..]
1151            .iter()
1152            .zip(leaf_products)
1153            .zip(points.chunks(block))
1154        {
1155            let mut remainder = Self::zeros(block);
1156            for (j, value) in values.iter().enumerate() {
1157                for (r, p) in remainder.data[..=j].iter_mut().zip(&product[block - j..]) {
1158                    *r += value.clone() * p;
1159                }
1160            }
1161            result.extend(points.iter().map(|p| remainder.eval(p.clone())));
1162        }
1163        result
1164    }
crates/competitive/src/math/number_theoretic_transform.rs (line 611)
581fn convolve_large_ntt<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
582where
583    M: Montgomery32NttModulus,
584{
585    let len = a.len() + b.len() - 1;
586    let ntt_len = 1usize << M::RANK;
587    let block_len = ntt_len / 2;
588    let same = a == b;
589    let transform = |a: &[MInt<M>]| {
590        let mut f = Vec::with_capacity(ntt_len);
591        advise_huge_pages(&mut f);
592        f.extend_from_slice(a);
593        Convolve::<M>::transform_ntt(f, ntt_len)
594    };
595    let fa: Vec<_> = a.chunks(block_len).map(transform).collect();
596    let fb: Option<Vec<_>> = if same {
597        None
598    } else {
599        Some(b.chunks(block_len).map(transform).collect())
600    };
601    let b_blocks = fb.as_ref().map_or(fa.len(), Vec::len);
602    let mut result = vec![MInt::<M>::zero(); len];
603    for diagonal in 0..fa.len() + b_blocks - 1 {
604        let mut spectrum = vec![MInt::<M>::zero(); ntt_len];
605        let start = diagonal.saturating_sub(b_blocks - 1);
606        for i in start..=diagonal.min(fa.len() - 1) {
607            let j = diagonal - i;
608            let g = if let Some(fb) = &fb { &fb[j] } else { &fa[j] };
609            pointwise_multiply_add(&mut spectrum, &fa[i], g);
610        }
611        spectrum = Convolve::<M>::inverse_transform_ntt(spectrum, ntt_len);
612        let offset = diagonal * block_len;
613        for (result, value) in result[offset..].iter_mut().zip(spectrum) {
614            *result += value;
615        }
616    }
617    result
618}
619
620impl<M> ConvolveSteps for Convolve<M>
621where
622    M: Montgomery32NttModulus,
623{
624    const CYCLIC: bool = true;
625
626    type T = Vec<MInt<M>>;
627    type F = Vec<MInt<M>>;
628    fn length(t: &Self::T) -> usize {
629        t.len()
630    }
631    fn transform(mut t: Self::T, len: usize) -> Self::F {
632        t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
633        #[cfg(target_arch = "x86_64")]
634        if use_block_ntt::<M>(t.len()) {
635            unsafe { ntt_simd::transform_blocks_avx2(&mut t) };
636            return t;
637        }
638        ntt(&mut t);
639        t
640    }
641    fn inverse_transform(mut f: Self::F, len: usize) -> Self::T {
642        #[cfg(target_arch = "x86_64")]
643        if use_block_ntt::<M>(f.len()) {
644            unsafe { ntt_simd::inverse_transform_blocks_avx2(&mut f) };
645            f.truncate(len);
646            return f;
647        }
648        intt(&mut f);
649        f.truncate(len);
650        f
651    }
652    fn multiply(f: &mut Self::F, g: &Self::F) {
653        assert_eq!(f.len(), g.len());
654        #[cfg(target_arch = "x86_64")]
655        if use_block_ntt::<M>(f.len()) {
656            unsafe { ntt_simd::multiply_blocks_avx2(f, g) };
657            return;
658        }
659        pointwise_multiply(f, g);
660    }
661    fn square(t: Self::T, len: usize) -> Self::T {
662        let mut f = Self::transform(t, len);
663        let g = f.clone();
664        Self::multiply(&mut f, &g);
665        Self::inverse_transform(f, len)
666    }
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    }
743}
744
745type MVec<M> = Vec<MInt<M>>;
746
747fn convert_crt_input<M, N1, N2, N3>(t: MVec<M>, capacity: usize) -> (MVec<N1>, MVec<N2>, MVec<N3>)
748where
749    M: MIntConvert<u32>,
750    N1: Montgomery32NttModulus,
751    N2: Montgomery32NttModulus,
752    N3: Montgomery32NttModulus,
753{
754    let mut f = (
755        MVec::<N1>::with_capacity(capacity),
756        MVec::<N2>::with_capacity(capacity),
757        MVec::<N3>::with_capacity(capacity),
758    );
759    advise_huge_pages(&mut f.0);
760    advise_huge_pages(&mut f.1);
761    advise_huge_pages(&mut f.2);
762    for t in t {
763        let t: u32 = t.into();
764        f.0.push(t.into());
765        f.1.push(t.into());
766        f.2.push(t.into());
767    }
768    f
769}
770
771fn reconstruct_mint_crt<M, N1, N2, N3>(f: (MVec<N1>, MVec<N2>, MVec<N3>)) -> MVec<M>
772where
773    M: MIntConvert + MIntConvert<u32>,
774    N1: Montgomery32NttModulus,
775    N2: Montgomery32NttModulus,
776    N3: Montgomery32NttModulus,
777{
778    let t1 = MInt::<N2>::new(N1::get_mod()).inv();
779    let m1_3 = MInt::<N3>::new(N1::get_mod());
780    let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
781    let modulus = <M as MIntConvert<u32>>::mod_into() as u64;
782    let m1 = N1::get_mod() as u64;
783    let m2 = m1 * N2::get_mod() as u64 % modulus;
784    let fits_u64 = (N1::get_mod() - 1) as u128
785        + (N2::get_mod() - 1) as u128 * m1 as u128
786        + (N3::get_mod() - 1) as u128 * m2 as u128
787        <= u64::MAX as u128;
788    f.0.into_iter()
789        .zip(f.1)
790        .zip(f.2)
791        .map(|((c1, c2), c3)| {
792            let d1 = c1.inner();
793            let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
794            let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
795            let d3 = ((c3 - x) * t2).inner();
796            let value = if fits_u64 {
797                (d1 as u64 + d2 as u64 * m1 + d3 as u64 * m2) % modulus
798            } else {
799                ((d1 as u128 + d2 as u128 * m1 as u128 + d3 as u128 * m2 as u128) % modulus as u128)
800                    as u64
801            };
802            MInt::<M>::from(value as u32)
803        })
804        .collect()
805}
806
807impl<M, N1, N2, N3> ConvolveSteps for Convolve<(M, (N1, N2, N3))>
808where
809    M: MIntConvert + MIntConvert<u32>,
810    N1: Montgomery32NttModulus,
811    N2: Montgomery32NttModulus,
812    N3: Montgomery32NttModulus,
813{
814    type T = MVec<M>;
815    type F = (MVec<N1>, MVec<N2>, MVec<N3>);
816    fn length(t: &Self::T) -> usize {
817        t.len()
818    }
819    fn transform(t: Self::T, len: usize) -> Self::F {
820        let npot = len.max(1).next_power_of_two();
821        let f = convert_crt_input(t, npot);
822        (
823            Convolve::<N1>::transform(f.0, npot),
824            Convolve::<N2>::transform(f.1, npot),
825            Convolve::<N3>::transform(f.2, npot),
826        )
827    }
828    fn inverse_transform(f: Self::F, len: usize) -> Self::T {
829        reconstruct_mint_crt((
830            Convolve::<N1>::inverse_transform(f.0, len),
831            Convolve::<N2>::inverse_transform(f.1, len),
832            Convolve::<N3>::inverse_transform(f.2, len),
833        ))
834    }
835    fn multiply(f: &mut Self::F, g: &Self::F) {
836        Convolve::<N1>::multiply(&mut f.0, &g.0);
837        Convolve::<N2>::multiply(&mut f.1, &g.1);
838        Convolve::<N3>::multiply(&mut f.2, &g.2);
839    }
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    }
881}
882
883fn convolve_mint_crt<M, N1, N2, N3>(a: MVec<M>, b: MVec<M>) -> MVec<M>
884where
885    M: MIntConvert + MIntConvert<u32>,
886    N1: Montgomery32NttModulus,
887    N2: Montgomery32NttModulus,
888    N3: Montgomery32NttModulus,
889{
890    let convolve = |a: MVec<M>, b: MVec<M>| {
891        let a_len = a.len();
892        let b_len = b.len();
893        let a = convert_crt_input(a, a_len);
894        let b = convert_crt_input(b, b_len);
895        reconstruct_mint_crt((
896            Convolve::<N1>::convolve(a.0, b.0),
897            Convolve::<N2>::convolve(a.1, b.1),
898            Convolve::<N3>::convolve(a.2, b.2),
899        ))
900    };
901    let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
902    let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
903    if a.len().min(b.len()) as u128 * (modulus - 1).pow(2) < capacity {
904        return convolve(a, b);
905    }
906    let block_len = ((capacity - 1) / (modulus - 1).pow(2)) as usize;
907    if block_len == 0 {
908        return convolve_naive(&a, &b);
909    }
910    let mut result = vec![MInt::<M>::zero(); a.len() + b.len() - 1];
911    for (i, a) in a.chunks(block_len).enumerate() {
912        for (j, b) in b.chunks(block_len).enumerate() {
913            let product = convolve(a.to_vec(), b.to_vec());
914            for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
915                *value += product;
916            }
917        }
918    }
919    result
920}
921
922impl<N1, N2, N3> ConvolveSteps for Convolve<(u64, (N1, N2, N3))>
923where
924    N1: Montgomery32NttModulus,
925    N2: Montgomery32NttModulus,
926    N3: Montgomery32NttModulus,
927{
928    type T = Vec<u64>;
929    type F = ([MVec<N1>; 3], [MVec<N2>; 3], [MVec<N3>; 3]);
930
931    fn length(t: &Self::T) -> usize {
932        t.len()
933    }
934
935    fn transform(t: Self::T, len: usize) -> Self::F {
936        let npot = len.max(1).next_power_of_two();
937        assert!(npot <= 1usize << N1::RANK.min(N2::RANK).min(N3::RANK));
938        // The 22-bit fallback needs room for three limb products per coefficient.
939        assert!(
940            3 * npot as u128 * ((1u128 << 22) - 1).pow(2)
941                < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
942        );
943        let bits = if 2 * npot as u128 * (u32::MAX as u128).pow(2)
944            < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
945        {
946            32
947        } else {
948            22
949        };
950        let parts = if bits == 32 && t.iter().all(|&value| value <= u32::MAX as u64) {
951            1
952        } else {
953            64usize.div_ceil(bits)
954        };
955        fn split<M: Montgomery32NttModulus>(
956            t: &[u64],
957            len: usize,
958            bits: usize,
959            parts: usize,
960        ) -> [MVec<M>; 3] {
961            std::array::from_fn(|part| {
962                if part >= parts {
963                    return Vec::new();
964                }
965                Convolve::<M>::transform(
966                    t.iter()
967                        .map(|&t| MInt::from((t >> (part * bits)) & ((1u64 << bits) - 1)))
968                        .collect(),
969                    len,
970                )
971            })
972        }
973        (
974            split(&t, npot, bits, parts),
975            split(&t, npot, bits, parts),
976            split(&t, npot, bits, parts),
977        )
978    }
979
980    fn inverse_transform(f: Self::F, len: usize) -> Self::T {
981        let bits = if f.0[2].is_empty() { 32 } else { 22 };
982        let t1 = MInt::<N2>::new(N1::get_mod()).inv();
983        let m1 = N1::get_mod() as u64;
984        let m1_3 = MInt::<N3>::new(N1::get_mod());
985        let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
986        let m2 = m1 * N2::get_mod() as u64;
987        let mut result = vec![0u64; len.min(f.0[0].len())];
988        for (part, ((f1, f2), f3)) in f.0.into_iter().zip(f.1).zip(f.2).enumerate() {
989            if f1.is_empty() {
990                continue;
991            }
992            for (value, ((c1, c2), c3)) in result.iter_mut().zip(
993                Convolve::<N1>::inverse_transform(f1, len)
994                    .into_iter()
995                    .zip(Convolve::<N2>::inverse_transform(f2, len))
996                    .zip(Convolve::<N3>::inverse_transform(f3, len)),
997            ) {
998                let d1 = c1.inner();
999                let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
1000                let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
1001                let d3 = ((c3 - x) * t2).inner();
1002                let limb = (d1 as u64)
1003                    .wrapping_add((d2 as u64).wrapping_mul(m1))
1004                    .wrapping_add((d3 as u64).wrapping_mul(m2));
1005                *value = value.wrapping_add(limb << (part * bits));
1006            }
1007        }
1008        result
1009    }
1010
1011    fn multiply(f: &mut Self::F, g: &Self::F) {
1012        fn multiply<M: Montgomery32NttModulus>(f: &mut [MVec<M>; 3], g: &[MVec<M>; 3]) {
1013            assert_eq!(f[0].len(), g[0].len());
1014            if f[1].is_empty() || g[1].is_empty() {
1015                if f[1].is_empty() && !g[1].is_empty() {
1016                    f[1] = f[0].clone();
1017                    Convolve::<M>::multiply(&mut f[1], &g[1]);
1018                } else if !f[1].is_empty() {
1019                    Convolve::<M>::multiply(&mut f[1], &g[0]);
1020                }
1021                Convolve::<M>::multiply(&mut f[0], &g[0]);
1022                return;
1023            }
1024            #[cfg(target_arch = "x86_64")]
1025            if use_block_ntt::<M>(f[0].len()) {
1026                for part in (1..if f[2].is_empty() { 2 } else { 3 }).rev() {
1027                    let mut sum = f[0].clone();
1028                    Convolve::<M>::multiply(&mut sum, &g[part]);
1029                    for left in 1..=part {
1030                        let mut product = f[left].clone();
1031                        Convolve::<M>::multiply(&mut product, &g[part - left]);
1032                        for (value, product) in sum.iter_mut().zip(product) {
1033                            // Block products contain lazy Montgomery residues.
1034                            *value = MInt::new(value.inner() + product.inner());
1035                        }
1036                    }
1037                    f[part] = sum;
1038                }
1039                Convolve::<M>::multiply(&mut f[0], &g[0]);
1040                return;
1041            }
1042            if f[2].is_empty() {
1043                for i in 0..f[0].len() {
1044                    f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1045                    f[0][i] *= g[0][i];
1046                }
1047                return;
1048            }
1049            for i in 0..f[0].len() {
1050                f[2][i] = f[0][i] * g[2][i] + f[1][i] * g[1][i] + f[2][i] * g[0][i];
1051                f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1052                f[0][i] *= g[0][i];
1053            }
1054        }
1055        multiply(&mut f.0, &g.0);
1056        multiply(&mut f.1, &g.1);
1057        multiply(&mut f.2, &g.2);
1058    }
1059
1060    fn square(t: Self::T, len: usize) -> Self::T {
1061        let mut f = Self::transform(t, len);
1062        let g = f.clone();
1063        Self::multiply(&mut f, &g);
1064        Self::inverse_transform(f, len)
1065    }
1066
1067    fn convolve(a: Self::T, b: Self::T) -> Self::T {
1068        let max_len = Self::length(&a).max(Self::length(&b));
1069        let min_len = Self::length(&a).min(Self::length(&b));
1070        let (balanced, short) = crate::avx_helper!(@dispatch_avx2_fma (300, 64), (1536, 512));
1071        if max_len <= balanced || min_len <= short {
1072            let a_wrapping: &[Wrapping<u64>] =
1073                unsafe { std::slice::from_raw_parts(a.as_ptr().cast(), a.len()) };
1074            let b_wrapping: &[Wrapping<u64>] =
1075                unsafe { std::slice::from_raw_parts(b.as_ptr().cast(), b.len()) };
1076            let mut c = std::mem::ManuallyDrop::new(if max_len <= 300 || min_len > 60 {
1077                convolve_karatsuba(a_wrapping, b_wrapping)
1078            } else {
1079                convolve_naive(a_wrapping, b_wrapping)
1080            });
1081            return unsafe { Vec::from_raw_parts(c.as_mut_ptr().cast(), c.len(), c.capacity()) };
1082        }
1083        let len = (Self::length(&a) + Self::length(&b)).saturating_sub(1);
1084        let block_len = if min_len >= 1 << 20 {
1085            1 << 20
1086        } else {
1087            (min_len.next_power_of_two() * 8).min(1 << 21) - min_len + 1
1088        };
1089        if max_len <= block_len {
1090            return convolve_u64_fft(a, b);
1091        }
1092        let mut result = vec![0u64; len];
1093        for (i, a) in a.chunks(block_len).enumerate() {
1094            for (j, b) in b.chunks(block_len).enumerate() {
1095                if a.len().min(b.len()) <= 60 {
1096                    for (x, &a) in a.iter().enumerate() {
1097                        for (y, &b) in b.iter().enumerate() {
1098                            let value = &mut result[(i + j) * block_len + x + y];
1099                            *value = value.wrapping_add(a.wrapping_mul(b));
1100                        }
1101                    }
1102                    continue;
1103                }
1104                let product = convolve_u64_fft(a.to_vec(), b.to_vec());
1105                for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
1106                    *value = value.wrapping_add(product);
1107                }
1108            }
1109        }
1110        result
1111    }
1112}
1113
1114fn convolve_u64_fft(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1115    // Keep limb convolutions below 2^47 at the 2^21 FFT limit.
1116    crate::avx_helper!(@dispatch_avx2_fma return unsafe {
1117        convolve_u64_avx2(a, b)
1118    }, ());
1119    convolve_u64_fft_scalar(a, b)
1120}
1121
1122fn convolve_u64_fft_scalar(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1123    fn split(values: &[u64]) -> [Vec<i64>; 5] {
1124        let mut result = std::array::from_fn(|_| Vec::with_capacity(values.len()));
1125        for mut value in values.iter().copied() {
1126            for part in &mut result {
1127                let digit = ((value << 51) as i64) >> 51;
1128                part.push(digit);
1129                value = (value >> 13).wrapping_add(u64::from(digit < 0));
1130            }
1131        }
1132        result
1133    }
1134
1135    let len = a.len() + b.len() - 1;
1136    let transform = |values: &[u64]| {
1137        if values.iter().any(|&value| value > u32::MAX as u64) {
1138            return split(values).map(|part| ConvolveRealFft::transform(part, len));
1139        }
1140        let [a, b, c, _, _] = split(values);
1141        let a = ConvolveRealFft::transform(a, len);
1142        let size = a.len();
1143        [
1144            a,
1145            ConvolveRealFft::transform(b, len),
1146            ConvolveRealFft::transform(c, len),
1147            vec![Zero::zero(); size],
1148            vec![Zero::zero(); size],
1149        ]
1150    };
1151    let fa = transform(&a);
1152    drop(a);
1153    let fb = transform(&b);
1154    drop(b);
1155    let values: [Vec<i64>; 5] = std::array::from_fn(|part| {
1156        let mut sum = fa[0].clone();
1157        ConvolveRealFft::multiply(&mut sum, &fb[part]);
1158        for left in 1..=part {
1159            let mut product = fa[left].clone();
1160            ConvolveRealFft::multiply(&mut product, &fb[part - left]);
1161            for (sum, product) in sum.iter_mut().zip(product) {
1162                *sum += product;
1163            }
1164        }
1165        ConvolveRealFft::inverse_transform(sum, len)
1166    });
1167    (0..len)
1168        .map(|i| {
1169            (values[0][i] as u64)
1170                .wrapping_add((values[1][i] as u64) << 13)
1171                .wrapping_add((values[2][i] as u64) << 26)
1172                .wrapping_add((values[3][i] as u64) << 39)
1173                .wrapping_add((values[4][i] as u64) << 52)
1174        })
1175        .collect()
1176}
1177
1178pub trait NttReuse: ConvolveSteps {
1179    const MULTIPLE: bool = true;
1180
1181    /// Transforms coefficients into the usual NTT frequency order.
1182    fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1183        Self::transform(t, len)
1184    }
1185
1186    /// Inverts a value produced by `transform_ntt`.
1187    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1188        Self::inverse_transform(f, len)
1189    }
1190
1191    /// Extends a value produced by `transform_ntt` to twice its length.
1192    /// If `monic`, the input represents a monic degree-`n` polynomial modulo
1193    /// `x^n - 1`, where `n` is the transform length.
1194    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F;
1195
1196    /// Extracts the even coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1197    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1198
1199    /// Extracts the odd coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1200    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1201
1202    /// Multiplies a usual NTT transform by the corresponding prefix of another one.
1203    fn multiply_prefix(f: &mut Self::F, g: &Self::F);
1204
1205    /// Adds the pointwise product of two usual NTT transforms to `sum`.
1206    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F);
1207
1208    /// Maximum number of products that can be summed before reconstruction.
1209    /// Both factors must transform canonical coefficients at the supplied transform's length,
1210    /// and each cyclic product must itself be reconstructible.
1211    fn max_product_sum_count(_f: &Self::F) -> usize {
1212        if Self::MULTIPLE { 1 } else { usize::MAX }
1213    }
1214
1215    fn power_projection_step(
1216        p_flat: Self::T,
1217        q_flat: Self::T,
1218        n: usize,
1219        py: usize,
1220        qy: usize,
1221    ) -> (Self::T, Self::T) {
1222        let base = n * 2;
1223        let len_p = base * py;
1224        let len_q = base * qy;
1225        let len = (len_p + len_q - 1).max(len_q + len_q - 1);
1226        let half = len.max(1).next_power_of_two() / 2;
1227
1228        let p_fft = Self::transform_ntt(p_flat, len);
1229        let q_fft = Self::transform_ntt(q_flat, len);
1230        let pr_fft = Self::odd_mul_normal_neg(&p_fft, &q_fft);
1231        let qr_fft = Self::even_mul_normal_neg(&q_fft, &q_fft);
1232        (
1233            Self::inverse_transform_ntt(pr_fft, half),
1234            Self::inverse_transform_ntt(qr_fft, half),
1235        )
1236    }
1237}
1238
1239thread_local!(
1240    static BIT_REVERSE: UnsafeCell<Vec<Vec<usize>>> = const { UnsafeCell::new(vec![]) };
1241);
1242
1243impl<M> NttReuse for Convolve<M>
1244where
1245    M: Montgomery32NttModulus,
1246{
1247    const MULTIPLE: bool = false;
1248
1249    fn transform_ntt(mut t: Self::T, len: usize) -> Self::F {
1250        t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
1251        ntt(&mut t);
1252        t
1253    }
1254
1255    fn inverse_transform_ntt(mut f: Self::F, len: usize) -> Self::T {
1256        intt(&mut f);
1257        f.truncate(len);
1258        f
1259    }
1260
1261    fn ntt_doubling(mut f: Self::F, monic: bool) -> Self::F {
1262        let n = f.len();
1263        let k = n.trailing_zeros() as usize;
1264        let mut a = Self::inverse_transform_ntt(f.clone(), n);
1265        if monic {
1266            a[0] -= MInt::<M>::from(2);
1267        }
1268        let zeta = MInt::<M>::new_unchecked(M::INFO.root[k + 1]);
1269        let zeta2 = zeta * zeta;
1270        let mut rot = [MInt::one(), zeta, zeta2, zeta2 * zeta];
1271        let step = zeta2 * zeta2;
1272        for a in a.chunks_mut(4) {
1273            for (a, rot) in a.iter_mut().zip(&mut rot) {
1274                *a *= *rot;
1275                *rot *= step;
1276            }
1277        }
1278        f.extend(Self::transform_ntt(a, n));
1279        f
1280    }
1281
1282    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1283        assert_eq!(f.len(), g.len());
1284        assert!(f.len().is_power_of_two());
1285        assert!(f.len() >= 2);
1286        if std::ptr::eq(f, g) {
1287            return f.as_chunks::<2>().0.iter().map(|a| a[0] * a[1]).collect();
1288        }
1289        let inv2 = MInt::<M>::from(2).inv();
1290        let n = f.len() / 2;
1291        (0..n)
1292            .map(|i| (f[i << 1] * g[i << 1 | 1] + f[i << 1 | 1] * g[i << 1]) * inv2)
1293            .collect()
1294    }
1295
1296    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1297        assert_eq!(f.len(), g.len());
1298        assert!(f.len().is_power_of_two());
1299        assert!(f.len() >= 2);
1300        let mut inv2 = MInt::<M>::from(2).inv();
1301        let n = f.len() / 2;
1302        let k = f.len().trailing_zeros() as usize;
1303        let mut h = vec![MInt::<M>::zero(); n];
1304        let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1305        BIT_REVERSE.with(|br| {
1306            let br = unsafe { &mut *br.get() };
1307            if br.len() < k {
1308                br.resize_with(k, Default::default);
1309            }
1310            let k = k - 1;
1311            if br[k].is_empty() {
1312                let mut v = vec![0; 1 << k];
1313                for i in 0..1 << k {
1314                    v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1315                }
1316                br[k] = v;
1317            }
1318            for &i in &br[k] {
1319                h[i] = (f[i << 1] * g[i << 1 | 1] - f[i << 1 | 1] * g[i << 1]) * inv2;
1320                inv2 *= w;
1321            }
1322        });
1323        h
1324    }
1325
1326    fn multiply_prefix(f: &mut Self::F, g: &Self::F) {
1327        pointwise_multiply(f, g);
1328    }
1329
1330    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F) {
1331        assert!(sum.len() == f.len() && sum.len() == g.len());
1332        pointwise_multiply_add(sum, f, g);
1333    }
1334
1335    fn power_projection_step(
1336        p_flat: Vec<MInt<M>>,
1337        q_flat: Vec<MInt<M>>,
1338        n: usize,
1339        py: usize,
1340        qy: usize,
1341    ) -> (Vec<MInt<M>>, Vec<MInt<M>>) {
1342        let high_degree = (qy - 1) * 2;
1343        let rows = (py + qy - 1).max(high_degree).next_power_of_two();
1344        let cols = n * 2;
1345        let size = rows * cols;
1346        let mut p = p_flat;
1347        p.resize_with(size, MInt::<M>::zero);
1348        ntt_rows(&mut p, cols);
1349        ntt_batch(&mut p, cols);
1350
1351        let mut q = q_flat;
1352        q.resize_with(size, MInt::<M>::zero);
1353        ntt_rows(&mut q, cols);
1354        let q_high = (rows == high_degree).then(|| q[(qy - 1) * cols..qy * cols].to_vec());
1355        ntt_batch(&mut q, cols);
1356
1357        let half = cols / 2;
1358        let mut odd_factor = vec![MInt::<M>::zero(); half];
1359        let mut factor = MInt::<M>::from(2).inv();
1360        let k = cols.trailing_zeros() as usize;
1361        let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1362        BIT_REVERSE.with(|br| {
1363            let br = unsafe { &mut *br.get() };
1364            if br.len() < k {
1365                br.resize_with(k, Default::default);
1366            }
1367            let k = k - 1;
1368            if br[k].is_empty() {
1369                let mut v = vec![0; 1 << k];
1370                for i in 0..1 << k {
1371                    v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1372                }
1373                br[k] = v;
1374            }
1375            for &i in &br[k] {
1376                odd_factor[i] = factor;
1377                factor *= w;
1378            }
1379        });
1380
1381        let mut pr = vec![MInt::<M>::zero(); rows * half];
1382        let mut qr = vec![MInt::<M>::zero(); rows * half];
1383        for i in 0..pr.len() {
1384            pr[i] = (p[i << 1] * q[i << 1 | 1] - p[i << 1 | 1] * q[i << 1])
1385                * odd_factor[i & (half - 1)];
1386            qr[i] = q[i << 1] * q[i << 1 | 1];
1387        }
1388        intt_batch(&mut pr, half);
1389        intt_rows(&mut pr, half);
1390        intt_batch(&mut qr, half);
1391        intt_rows(&mut qr, half);
1392
1393        if let Some(q_high) = q_high {
1394            let mut q_high_even = vec![MInt::<M>::zero(); half];
1395            for i in 0..half {
1396                q_high_even[i] = q_high[i << 1] * q_high[i << 1 | 1];
1397            }
1398            intt(&mut q_high_even);
1399            for (value, high) in qr.iter_mut().zip(&q_high_even) {
1400                *value -= *high;
1401            }
1402            qr.extend_from_slice(&q_high_even);
1403        }
1404        (pr, qr)
1405    }
1406}
1407
1408impl<M, N1, N2, N3> NttReuse for Convolve<(M, (N1, N2, N3))>
1409where
1410    M: MIntConvert + MIntConvert<u32>,
1411    N1: Montgomery32NttModulus,
1412    N2: Montgomery32NttModulus,
1413    N3: Montgomery32NttModulus,
1414{
1415    fn max_product_sum_count(f: &Self::F) -> usize {
1416        let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
1417        if modulus == 1 {
1418            return usize::MAX;
1419        }
1420        let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
1421        ((capacity - 1) / ((modulus - 1) * (modulus - 1)) / f.0.len() as u128)
1422            .clamp(1, usize::MAX as u128) as usize
1423    }
1424
1425    fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1426        let npot = len.max(1).next_power_of_two();
1427        let f = convert_crt_input(t, npot);
1428        (
1429            Convolve::<N1>::transform_ntt(f.0, npot),
1430            Convolve::<N2>::transform_ntt(f.1, npot),
1431            Convolve::<N3>::transform_ntt(f.2, npot),
1432        )
1433    }
1434
1435    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1436        reconstruct_mint_crt((
1437            Convolve::<N1>::inverse_transform_ntt(f.0, len),
1438            Convolve::<N2>::inverse_transform_ntt(f.1, len),
1439            Convolve::<N3>::inverse_transform_ntt(f.2, len),
1440        ))
1441    }
1442
1443    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F {
1444        if monic {
1445            let n = f.0.len();
1446            let mut coefficients = Self::inverse_transform_ntt(f, n);
1447            coefficients[0] -= MInt::<M>::one();
1448            coefficients.push(MInt::<M>::one());
1449            Self::transform_ntt(coefficients, n * 2)
1450        } else {
1451            (
1452                Convolve::<N1>::ntt_doubling(f.0, false),
1453                Convolve::<N2>::ntt_doubling(f.1, false),
1454                Convolve::<N3>::ntt_doubling(f.2, false),
1455            )
1456        }
1457    }
Source

fn max_product_sum_count(_f: &Self::F) -> usize

Maximum number of products that can be summed before reconstruction. Both factors must transform canonical coefficients at the supplied transform’s length, and each cyclic product must itself be reconstructible.

Examples found in repository?
crates/competitive/src/math/formal_power_series/formal_power_series_impls.rs (line 378)
373    fn sum_products(f: &[C::F], g: &[C::F], len: usize) -> Vec<T>
374    where
375        C: NttReuse<T = Vec<T>>,
376        C::F: Clone,
377    {
378        let chunk = C::max_product_sum_count(&f[0]);
379        f.rchunks(chunk)
380            .zip(g.chunks(chunk))
381            .map(|(f, g)| {
382                let mut sum = f[f.len() - 1].clone();
383                C::multiply_prefix(&mut sum, &g[0]);
384                for (f, g) in f.iter().rev().skip(1).zip(&g[1..]) {
385                    C::multiply_add(&mut sum, f, g);
386                }
387                C::inverse_transform_ntt(sum, len)
388            })
389            .reduce(|mut sum, part| {
390                for (sum, value) in sum.iter_mut().zip(part) {
391                    *sum += value;
392                }
393                sum
394            })
395            .unwrap()
396    }
Source

fn power_projection_step( p_flat: Self::T, q_flat: Self::T, n: usize, py: usize, qy: usize, ) -> (Self::T, Self::T)

Examples found in repository?
crates/competitive/src/math/formal_power_series/formal_power_series_impls.rs (line 1290)
1258    pub fn power_projection(&self, w: &[T], m: usize) -> Self
1259    where
1260        C: NttReuse<T = Vec<T>>,
1261    {
1262        if w.is_empty() {
1263            return Self::zeros(m);
1264        }
1265        if m <= 1 {
1266            return Self::from_vec(vec![w[0].clone(); m]);
1267        }
1268
1269        let n0 = w.len();
1270        let mut n = n0.next_power_of_two();
1271        let mut f = self.prefix_ref(n);
1272        f.resize(n);
1273
1274        let base = n * 2;
1275        let mut p_flat = vec![T::zero(); base];
1276        for (i, wi) in w.iter().enumerate() {
1277            p_flat[n - 1 - i] = wi.clone();
1278        }
1279        let mut q_flat = vec![T::zero(); base * 2];
1280        q_flat[0] = T::one();
1281        let q_offset = base;
1282        for (i, fi) in f.iter().enumerate() {
1283            q_flat[q_offset + i] = -fi.clone();
1284        }
1285        let mut py = 1usize;
1286        let mut qy = 2usize;
1287
1288        let y_limit = m;
1289        while n > 1 {
1290            let (mut p, mut q) = C::power_projection_step(p_flat, q_flat, n, py, qy);
1291            let new_py = (py + qy - 1).min(y_limit);
1292            let new_qy = (qy + qy - 1).min(y_limit);
1293            p.resize_with(n * new_py, T::zero);
1294            q.resize_with(n * new_qy, T::zero);
1295
1296            let n2 = n / 2;
1297            for row in p.chunks_exact_mut(n) {
1298                row[n2..].fill_with(T::zero);
1299            }
1300            for row in q.chunks_exact_mut(n) {
1301                row[n2..].fill_with(T::zero);
1302            }
1303            p_flat = p;
1304            q_flat = q;
1305            py = new_py;
1306            qy = new_qy;
1307            n = n2;
1308        }
1309
1310        let base = 2;
1311        let mut p_y = Vec::with_capacity(py);
1312        for y in 0..py {
1313            p_y.push(p_flat[base * y].clone());
1314        }
1315        let mut q_y = Vec::with_capacity(qy);
1316        for y in 0..qy {
1317            q_y.push(q_flat[base * y].clone());
1318        }
1319        (Self::from_vec(p_y) * Self::from_vec(q_y).inv(m)).prefix(m)
1320    }

Dyn Compatibility§

This trait is not dyn compatible.

In older versions of Rust, dyn compatibility was called "object safety".

Implementors§