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§
Required Methods§
Sourcefn ntt_doubling(f: Self::F, monic: bool) -> Self::F
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.
Sourcefn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F
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.
Sourcefn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F
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.
Sourcefn multiply_prefix(f: &mut Self::F, g: &Self::F)
fn multiply_prefix(f: &mut Self::F, g: &Self::F)
Multiplies a usual NTT transform by the corresponding prefix of another one.
Sourcefn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F)
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§
Sourcefn transform_ntt(t: Self::T, len: usize) -> Self::F
fn transform_ntt(t: Self::T, len: usize) -> Self::F
Transforms coefficients into the usual NTT frequency order.
Examples found in repository?
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
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 }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, "ient);
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 }Sourcefn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T
fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T
Inverts a value produced by transform_ntt.
Examples found in repository?
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
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, "ient);
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 }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 }Sourcefn max_product_sum_count(_f: &Self::F) -> usize
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?
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 }Sourcefn power_projection_step(
p_flat: Self::T,
q_flat: Self::T,
n: usize,
py: usize,
qy: usize,
) -> (Self::T, Self::T)
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?
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".