#[repr(transparent)]pub struct MInt<M>where
M: MIntBase,{
x: M::Inner,
_marker: PhantomData<fn() -> M>,
}Fields§
§x: M::Inner§_marker: PhantomData<fn() -> M>Implementations§
Source§impl<M> MInt<M>where
M: MIntConvert<u32>,
impl<M> MInt<M>where
M: MIntConvert<u32>,
Sourcepub fn sqrt(self) -> Option<Self>
pub fn sqrt(self) -> Option<Self>
Examples found in repository?
More examples
crates/library_checker/src/number_theory/sqrt_mod.rs (line 10)
5pub fn sqrt_mod(reader: impl Read, writer: impl Write) {
6 prepare_io!(reader, writer);
7 sc!(q, yp: [(u32, u32); iter q]);
8 for (y, p) in yp {
9 DynMIntU32::set_mod(p);
10 if let Some(x) = DynMIntU32::from(y).sqrt() {
11 pp!(x);
12 } else {
13 pp!("-1");
14 }
15 }
16}Source§impl<M> MInt<M>where
M: MIntConvert,
impl<M> MInt<M>where
M: MIntConvert,
Sourcepub fn new(x: M::Inner) -> Self
pub fn new(x: M::Inner) -> Self
Examples found in repository?
crates/library_checker/src/number_theory/sum_of_totient_function.rs (line 15)
9pub fn sum_of_totient_function(reader: impl Read, writer: impl Write) {
10 prepare_io!(reader, writer);
11 sc!(n: u64);
12 let mut s = 1;
13 let mut pp = 0;
14 let mut pc = 0;
15 let inv2 = M::new(2).inv();
16 let qa = QuotientArray::from_fn(n, |i| [M::from(i), M::from(i) * M::from(i + 1) * inv2])
17 .map(|[x, y]| [x - M::one(), y - M::one()])
18 .lucy_dp::<ArrayOperation<AdditiveOperation<_>, 2>>(|[x, y], p| [x, y * M::from(p)])
19 .map(|[x, y]| y - x)
20 .min_25_sieve::<AddMulOperation<_>>(|p, c| {
21 if pp != p || pc > c {
22 pp = p;
23 pc = 1;
24 s = p - 1;
25 }
26 while pc < c {
27 pc += 1;
28 s *= p;
29 }
30 M::from(s)
31 });
32 pp!(qa[n]);
33}More examples
crates/competitive/src/math/number_theoretic_transform.rs (line 778)
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 }
1458
1459 fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1460 fn even_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1461 where
1462 M: Montgomery32NttModulus,
1463 {
1464 let n = f.len();
1465 assert_eq!(f.len(), g.len());
1466 assert!(f.len().is_power_of_two());
1467 assert!(f.len() >= 2);
1468 let inv2 = MInt::<M>::from(2).inv();
1469 let u = MInt::<M>::new(m) * MInt::<M>::from(n as u32);
1470 let n = f.len() / 2;
1471 (0..n)
1472 .map(|i| {
1473 (f[i << 1]
1474 * if i == 0 {
1475 g[i << 1 | 1] + u
1476 } else {
1477 g[i << 1 | 1]
1478 }
1479 + f[i << 1 | 1] * g[i << 1])
1480 * inv2
1481 })
1482 .collect()
1483 }
1484
1485 let m = M::mod_into();
1486 (
1487 even_mul_normal_neg_corrected(&f.0, &g.0, m),
1488 even_mul_normal_neg_corrected(&f.1, &g.1, m),
1489 even_mul_normal_neg_corrected(&f.2, &g.2, m),
1490 )
1491 }
1492
1493 fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1494 fn odd_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1495 where
1496 M: Montgomery32NttModulus,
1497 {
1498 assert_eq!(f.len(), g.len());
1499 assert!(f.len().is_power_of_two());
1500 assert!(f.len() >= 2);
1501 let mut inv2 = MInt::<M>::from(2).inv();
1502 let u = MInt::<M>::new(m) * MInt::<M>::from(f.len() as u32);
1503 let n = f.len() / 2;
1504 let k = f.len().trailing_zeros() as usize;
1505 let mut h = vec![MInt::<M>::zero(); n];
1506 let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1507 BIT_REVERSE.with(|br| {
1508 let br = unsafe { &mut *br.get() };
1509 if br.len() < k {
1510 br.resize_with(k, Default::default);
1511 }
1512 let k = k - 1;
1513 if br[k].is_empty() {
1514 let mut v = vec![0; 1 << k];
1515 for i in 0..1 << k {
1516 v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1517 }
1518 br[k] = v;
1519 }
1520 for &i in &br[k] {
1521 h[i] = (f[i << 1]
1522 * if i == 0 {
1523 g[i << 1 | 1] + u
1524 } else {
1525 g[i << 1 | 1]
1526 }
1527 - f[i << 1 | 1] * g[i << 1])
1528 * inv2;
1529 inv2 *= w;
1530 }
1531 });
1532 h
1533 }Source§impl<M> MInt<M>where
M: MIntBase,
impl<M> MInt<M>where
M: MIntBase,
Sourcepub const fn new_unchecked(x: M::Inner) -> Self
pub const fn new_unchecked(x: M::Inner) -> Self
Examples found in repository?
crates/competitive/src/num/mint/mint_base.rs (line 60)
59 pub fn new(x: M::Inner) -> Self {
60 Self::new_unchecked(<M as MIntConvert<M::Inner>>::from(x))
61 }
62}
63impl<M> MInt<M>
64where
65 M: MIntBase,
66{
67 #[inline]
68 pub const fn new_unchecked(x: M::Inner) -> Self {
69 Self {
70 x,
71 _marker: PhantomData,
72 }
73 }
74 #[inline]
75 pub fn get_mod() -> M::Inner {
76 M::get_mod()
77 }
78 #[inline]
79 pub fn pow(self, y: usize) -> Self {
80 Self::new_unchecked(M::mod_pow(self.x, y))
81 }
82 #[inline]
83 pub fn inv(self) -> Self {
84 Self::new_unchecked(M::mod_inv(self.x))
85 }
86 #[inline]
87 pub fn inner(self) -> M::Inner {
88 M::mod_inner(self.x)
89 }
90}
91
92impl<M> Clone for MInt<M>
93where
94 M: MIntBase,
95{
96 #[inline]
97 fn clone(&self) -> Self {
98 *self
99 }
100}
101impl<M> Copy for MInt<M> where M: MIntBase {}
102impl<M> Debug for MInt<M>
103where
104 M: MIntBase,
105{
106 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
107 Debug::fmt(&self.inner(), f)
108 }
109}
110impl<M> Default for MInt<M>
111where
112 M: MIntBase,
113{
114 #[inline]
115 fn default() -> Self {
116 <Self as Zero>::zero()
117 }
118}
119impl<M> PartialEq for MInt<M>
120where
121 M: MIntBase,
122{
123 #[inline]
124 fn eq(&self, other: &Self) -> bool {
125 PartialEq::eq(&self.x, &other.x)
126 }
127}
128impl<M> Eq for MInt<M> where M: MIntBase {}
129impl<M> Hash for MInt<M>
130where
131 M: MIntBase,
132{
133 #[inline]
134 fn hash<H: Hasher>(&self, state: &mut H) {
135 Hash::hash(&self.x, state)
136 }
137}
138macro_rules! impl_mint_from {
139 ($($t:ty),*) => {
140 $(impl<M> From<$t> for MInt<M>
141 where
142 M: MIntConvert<$t>,
143 {
144 #[inline]
145 fn from(x: $t) -> Self {
146 Self::new_unchecked(<M as MIntConvert<$t>>::from(x))
147 }
148 }
149 impl<M> From<MInt<M>> for $t
150 where
151 M: MIntConvert<$t>,
152 {
153 #[inline]
154 fn from(x: MInt<M>) -> $t {
155 <M as MIntConvert<$t>>::into(x.x)
156 }
157 })*
158 };
159}
160impl_mint_from!(
161 u8, u16, u32, u64, u128, usize, i8, i16, i32, i64, i128, isize
162);
163impl<M> Zero for MInt<M>
164where
165 M: MIntBase,
166{
167 #[inline]
168 fn zero() -> Self {
169 Self::new_unchecked(M::mod_zero())
170 }
171}
172impl<M> One for MInt<M>
173where
174 M: MIntBase,
175{
176 #[inline]
177 fn one() -> Self {
178 Self::new_unchecked(M::mod_one())
179 }
180}
181
182impl<M> Add for MInt<M>
183where
184 M: MIntBase,
185{
186 type Output = Self;
187 #[inline]
188 fn add(self, rhs: Self) -> Self::Output {
189 Self::new_unchecked(M::mod_add(self.x, rhs.x))
190 }
191}
192impl<M> Sub for MInt<M>
193where
194 M: MIntBase,
195{
196 type Output = Self;
197 #[inline]
198 fn sub(self, rhs: Self) -> Self::Output {
199 Self::new_unchecked(M::mod_sub(self.x, rhs.x))
200 }
201}
202impl<M> Mul for MInt<M>
203where
204 M: MIntBase,
205{
206 type Output = Self;
207 #[inline]
208 fn mul(self, rhs: Self) -> Self::Output {
209 Self::new_unchecked(M::mod_mul(self.x, rhs.x))
210 }
211}
212impl<M> Div for MInt<M>
213where
214 M: MIntBase,
215{
216 type Output = Self;
217 #[inline]
218 fn div(self, rhs: Self) -> Self::Output {
219 Self::new_unchecked(M::mod_div(self.x, rhs.x))
220 }
221}
222impl<M> Neg for MInt<M>
223where
224 M: MIntBase,
225{
226 type Output = Self;
227 #[inline]
228 fn neg(self) -> Self::Output {
229 Self::new_unchecked(M::mod_neg(self.x))
230 }More examples
crates/competitive/src/num/mint/mod.rs (line 81)
77 fn deserialize<I>(iter: &mut I) -> Self
78 where
79 I: Iterator<Item = u8>,
80 {
81 Self::new_unchecked(M::Inner::deserialize(iter))
82 }
83}
84
85#[cfg_attr(nightly, codesnip::entry(when("MIntBase", "random_generator")))]
86mod random_spec {
87 use super::*;
88 use std::ops::{RangeFull, RangeTo};
89
90 impl<M> RandomSpec<MInt<M>> for RangeFull
91 where
92 M: MIntBase,
93 RangeTo<M::Inner>: RandomSpec<M::Inner>,
94 {
95 fn rand(&self, rng: &mut Xorshift) -> MInt<M> {
96 MInt::<M>::new_unchecked(rng.random(..M::get_mod()))
97 }crates/library_checker/src/tree/point_set_tree_path_composite_sum.rs (line 50)
46 fn operate(x: &Self::T, y: &Self::T) -> Self::T {
47 Path {
48 a: x.a * y.a,
49 b: x.b + x.a * y.b,
50 sum: x.sum + x.a * y.sum + x.b * M::new_unchecked(y.cnt),
51 cnt: x.cnt + y.cnt,
52 }
53 }
54}
55impl Unital for PathMonoid {
56 fn unit() -> Self::T {
57 Path {
58 a: M::one(),
59 b: M::zero(),
60 sum: M::zero(),
61 cnt: 0,
62 }
63 }
64}
65impl Associative for PathMonoid {}
66
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68struct PathPair {
69 forward: Path,
70 reverse: Path,
71}
72
73struct PathPairMonoid;
74impl Magma for PathPairMonoid {
75 type T = PathPair;
76 fn operate(x: &Self::T, y: &Self::T) -> Self::T {
77 PathPair {
78 forward: PathMonoid::operate(&x.forward, &y.forward),
79 reverse: PathMonoid::operate(&y.reverse, &x.reverse),
80 }
81 }
82}
83impl Unital for PathPairMonoid {
84 fn unit() -> Self::T {
85 PathPair {
86 forward: PathMonoid::unit(),
87 reverse: PathMonoid::unit(),
88 }
89 }
90}
91impl Associative for PathPairMonoid {}
92
93struct Dp;
94
95impl MonoidCluster for Dp {
96 type Vertex = M;
97 type Edge = (M, M);
98 type PointMonoid = PointMonoid;
99 type PathMonoid = PathPairMonoid;
100
101 fn add_vertex(point: &Point, vertex: &M, parent_edge: Option<&(M, M)>) -> PathPair {
102 let cnt = point.cnt + 1;
103 let subtotal = point.sum + *vertex;
104 let (a, b) = parent_edge.copied().unwrap_or((M::one(), M::zero()));
105 PathPair {
106 forward: Path {
107 a,
108 b,
109 sum: a * subtotal + b * M::new_unchecked(cnt),
110 cnt,
111 },
112 reverse: Path {
113 a,
114 b,
115 sum: subtotal,
116 cnt,
117 },
118 }
119 }crates/library_checker/src/tree/point_set_tree_path_composite_sum_fixed_root.rs (line 50)
46 fn operate(x: &Self::T, y: &Self::T) -> Self::T {
47 Path {
48 a: x.a * y.a,
49 b: x.b + x.a * y.b,
50 sum: x.sum + x.a * y.sum + x.b * M::new_unchecked(y.cnt),
51 cnt: x.cnt + y.cnt,
52 }
53 }
54}
55impl Unital for PathMonoid {
56 fn unit() -> Self::T {
57 Path {
58 a: M::one(),
59 b: M::zero(),
60 sum: M::zero(),
61 cnt: 0,
62 }
63 }
64}
65impl Associative for PathMonoid {}
66
67struct Dp;
68
69impl MonoidCluster for Dp {
70 type Vertex = M;
71 type Edge = (M, M);
72 type PointMonoid = PointMonoid;
73 type PathMonoid = PathMonoid;
74
75 fn add_vertex(point: &Point, vertex: &M, parent_edge: Option<&(M, M)>) -> Path {
76 let cnt = point.cnt + 1;
77 let subtotal = point.sum + *vertex;
78 let (a, b) = parent_edge.copied().unwrap_or((M::one(), M::zero()));
79 Path {
80 a,
81 b,
82 sum: a * subtotal + b * M::new_unchecked(cnt),
83 cnt,
84 }
85 }crates/library_checker/src/tree/tree_path_composite_sum.rs (line 18)
8pub fn tree_path_composite_sum(reader: impl Read, writer: impl Write) {
9 prepare_io!(reader, writer);
10 sc!(n, values: [M; n], (graph, edges): @TreeGraphScanner::<usize, (M, M)>::new(n));
11 let dp = ReRooting::<(AdditiveOperation<M>, AdditiveOperation<i32>), _>::new_with_inverse(
12 &graph,
13 |&(sum, count), v, edge| {
14 let sum = sum + values[v];
15 let count = count + 1;
16 if let Some(edge) = edge {
17 let (a, b) = edges[edge];
18 (a * sum + b * M::new_unchecked(count as u32), count)
19 } else {
20 (sum, count)
21 }
22 },
23 );
24 pp!(@it dp.dp.iter().map(|value| value.0));
25}pub fn get_mod() -> M::Inner
Sourcepub fn pow(self, y: usize) -> Self
pub fn pow(self, y: usize) -> Self
Examples found in repository?
More examples
crates/competitive/src/math/mint_matrix.rs (line 237)
234 fn x_pow_mod(&self, k: usize) -> Self {
235 let d = self.0.len() - 1;
236 if d == 1 {
237 return Self(vec![(-self.0[0]).pow(k)]);
238 }
239 let mut r = Self(vec![MInt::zero(); d]);
240 r.0[0] = MInt::one();
241 for bit in (0..usize::BITS - k.leading_zeros()).rev() {
242 r = r.square_mod(self);
243 if k >> bit & 1 != 0 {
244 let x = r.0[d - 1];
245 for i in (1..d).rev() {
246 r.0[i] = r.0[i - 1] - x * self.0[i];
247 }
248 r.0[0] = -x * self.0[0];
249 }
250 }
251 r
252 }Sourcepub fn inv(self) -> Self
pub fn inv(self) -> Self
Examples found in repository?
crates/competitive/src/math/factorial.rs (line 22)
16 pub fn new(max_n: usize) -> Self {
17 let mut fact = vec![MInt::one(); max_n + 1];
18 let mut inv_fact = vec![MInt::one(); max_n + 1];
19 for i in 2..=max_n {
20 fact[i] = fact[i - 1] * MInt::from(i);
21 }
22 inv_fact[max_n] = fact[max_n].inv();
23 for i in (3..=max_n).rev() {
24 inv_fact[i - 1] = inv_fact[i] * MInt::from(i);
25 }
26 Self { fact, inv_fact }
27 }More examples
crates/competitive/src/math/black_box_mint_matrix.rs (line 80)
72 fn black_box_linear_equation(&self, mut b: Vec<MInt<M>>) -> Option<Vec<MInt<M>>> {
73 assert_eq!(self.shape().0, self.shape().1);
74 assert_eq!(self.shape().1, b.len());
75 let n = self.shape().0;
76 let p = self.minimal_polynomial();
77 if p.is_empty() || p[0].is_zero() {
78 return None;
79 }
80 let p0_inv = p[0].inv();
81 let mut x = vec![MInt::zero(); n];
82 for p in p.into_iter().skip(1) {
83 let p = -p * p0_inv;
84 for i in 0..n {
85 x[i] += p * b[i];
86 }
87 b = self.apply(&b);
88 }
89 Some(x)
90 }crates/library_checker/src/number_theory/sum_of_totient_function.rs (line 15)
9pub fn sum_of_totient_function(reader: impl Read, writer: impl Write) {
10 prepare_io!(reader, writer);
11 sc!(n: u64);
12 let mut s = 1;
13 let mut pp = 0;
14 let mut pc = 0;
15 let inv2 = M::new(2).inv();
16 let qa = QuotientArray::from_fn(n, |i| [M::from(i), M::from(i) * M::from(i + 1) * inv2])
17 .map(|[x, y]| [x - M::one(), y - M::one()])
18 .lucy_dp::<ArrayOperation<AdditiveOperation<_>, 2>>(|[x, y], p| [x, y * M::from(p)])
19 .map(|[x, y]| y - x)
20 .min_25_sieve::<AddMulOperation<_>>(|p, c| {
21 if pp != p || pc > c {
22 pp = p;
23 pc = 1;
24 s = p - 1;
25 }
26 while pc < c {
27 pc += 1;
28 s *= p;
29 }
30 M::from(s)
31 });
32 pp!(qa[n]);
33}crates/competitive/src/math/binomial_prefix_sum.rs (line 75)
60 pub fn for_each<F>(self, mut f: F)
61 where
62 F: FnMut(usize, MInt<M>),
63 {
64 let query = &self.query;
65 if query.is_empty() {
66 return;
67 }
68 let max_n = query.iter().map(|&(_, n)| n).max().unwrap_or(0);
69 let modulus = M::mod_into();
70 debug_assert!(modulus > 2 && modulus % 2 == 1);
71 debug_assert!(max_n < modulus);
72 debug_assert!(query.iter().all(|&(m, n)| m <= n));
73
74 let fact = MemorizedFactorial::<M>::new(max_n);
75 let inv2 = MInt::<M>::from(2usize).inv();
76 let mut cur = MInt::<M>::one();
77 crate::mo_algorithm!(
78 query,
79 (m, n),
80 |old_m| cur += fact.combination(n, old_m + 1),
81 |new_m| cur -= fact.combination(n, new_m + 1),
82 |old_n| cur = cur + cur - fact.combination(old_n, m),
83 |new_n| cur = (cur + fact.combination(new_n, m)) * inv2,
84 |i| f(i, cur),
85 );
86 }crates/competitive/src/math/lagrange_interpolation.rs (line 78)
50pub fn lagrange_interpolation_polynomial<M>(x: &[MInt<M>], y: &[MInt<M>]) -> Vec<MInt<M>>
51where
52 M: MIntBase,
53{
54 let n = x.len() - 1;
55 let mut dp = vec![MInt::zero(); n + 2];
56 let mut ndp = vec![MInt::zero(); n + 2];
57 dp[0] = -x[0];
58 dp[1] = MInt::one();
59 for x in x.iter().skip(1) {
60 for j in 0..=n + 1 {
61 ndp[j] = -dp[j] * x + if j >= 1 { dp[j - 1] } else { MInt::zero() };
62 }
63 std::mem::swap(&mut dp, &mut ndp);
64 }
65 let mut res = vec![MInt::zero(); n + 1];
66 for i in 0..=n {
67 let t = y[i]
68 / (0..=n)
69 .map(|j| if i != j { x[i] - x[j] } else { MInt::one() })
70 .product::<MInt<M>>();
71 if t.is_zero() {
72 continue;
73 } else if x[i].is_zero() {
74 for j in 0..=n {
75 res[j] += dp[j + 1] * t;
76 }
77 } else {
78 let xinv = x[i].inv();
79 let mut pre = MInt::zero();
80 for j in 0..=n {
81 let d = -(dp[j] - pre) * xinv;
82 res[j] += d * t;
83 pre = d;
84 }
85 }
86 }
87 res
88}crates/competitive/src/math/number_theoretic_transform.rs (line 778)
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 }
1458
1459 fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1460 fn even_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1461 where
1462 M: Montgomery32NttModulus,
1463 {
1464 let n = f.len();
1465 assert_eq!(f.len(), g.len());
1466 assert!(f.len().is_power_of_two());
1467 assert!(f.len() >= 2);
1468 let inv2 = MInt::<M>::from(2).inv();
1469 let u = MInt::<M>::new(m) * MInt::<M>::from(n as u32);
1470 let n = f.len() / 2;
1471 (0..n)
1472 .map(|i| {
1473 (f[i << 1]
1474 * if i == 0 {
1475 g[i << 1 | 1] + u
1476 } else {
1477 g[i << 1 | 1]
1478 }
1479 + f[i << 1 | 1] * g[i << 1])
1480 * inv2
1481 })
1482 .collect()
1483 }
1484
1485 let m = M::mod_into();
1486 (
1487 even_mul_normal_neg_corrected(&f.0, &g.0, m),
1488 even_mul_normal_neg_corrected(&f.1, &g.1, m),
1489 even_mul_normal_neg_corrected(&f.2, &g.2, m),
1490 )
1491 }
1492
1493 fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1494 fn odd_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1495 where
1496 M: Montgomery32NttModulus,
1497 {
1498 assert_eq!(f.len(), g.len());
1499 assert!(f.len().is_power_of_two());
1500 assert!(f.len() >= 2);
1501 let mut inv2 = MInt::<M>::from(2).inv();
1502 let u = MInt::<M>::new(m) * MInt::<M>::from(f.len() as u32);
1503 let n = f.len() / 2;
1504 let k = f.len().trailing_zeros() as usize;
1505 let mut h = vec![MInt::<M>::zero(); n];
1506 let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1507 BIT_REVERSE.with(|br| {
1508 let br = unsafe { &mut *br.get() };
1509 if br.len() < k {
1510 br.resize_with(k, Default::default);
1511 }
1512 let k = k - 1;
1513 if br[k].is_empty() {
1514 let mut v = vec![0; 1 << k];
1515 for i in 0..1 << k {
1516 v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1517 }
1518 br[k] = v;
1519 }
1520 for &i in &br[k] {
1521 h[i] = (f[i << 1]
1522 * if i == 0 {
1523 g[i << 1 | 1] + u
1524 } else {
1525 g[i << 1 | 1]
1526 }
1527 - f[i << 1 | 1] * g[i << 1])
1528 * inv2;
1529 inv2 *= w;
1530 }
1531 });
1532 h
1533 }Additional examples can be found in:
Sourcepub fn inner(self) -> M::Inner
pub fn inner(self) -> M::Inner
Examples found in repository?
crates/competitive/src/num/mint/mint_base.rs (line 107)
106 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
107 Debug::fmt(&self.inner(), f)
108 }
109}
110impl<M> Default for MInt<M>
111where
112 M: MIntBase,
113{
114 #[inline]
115 fn default() -> Self {
116 <Self as Zero>::zero()
117 }
118}
119impl<M> PartialEq for MInt<M>
120where
121 M: MIntBase,
122{
123 #[inline]
124 fn eq(&self, other: &Self) -> bool {
125 PartialEq::eq(&self.x, &other.x)
126 }
127}
128impl<M> Eq for MInt<M> where M: MIntBase {}
129impl<M> Hash for MInt<M>
130where
131 M: MIntBase,
132{
133 #[inline]
134 fn hash<H: Hasher>(&self, state: &mut H) {
135 Hash::hash(&self.x, state)
136 }
137}
138macro_rules! impl_mint_from {
139 ($($t:ty),*) => {
140 $(impl<M> From<$t> for MInt<M>
141 where
142 M: MIntConvert<$t>,
143 {
144 #[inline]
145 fn from(x: $t) -> Self {
146 Self::new_unchecked(<M as MIntConvert<$t>>::from(x))
147 }
148 }
149 impl<M> From<MInt<M>> for $t
150 where
151 M: MIntConvert<$t>,
152 {
153 #[inline]
154 fn from(x: MInt<M>) -> $t {
155 <M as MIntConvert<$t>>::into(x.x)
156 }
157 })*
158 };
159}
160impl_mint_from!(
161 u8, u16, u32, u64, u128, usize, i8, i16, i32, i64, i128, isize
162);
163impl<M> Zero for MInt<M>
164where
165 M: MIntBase,
166{
167 #[inline]
168 fn zero() -> Self {
169 Self::new_unchecked(M::mod_zero())
170 }
171}
172impl<M> One for MInt<M>
173where
174 M: MIntBase,
175{
176 #[inline]
177 fn one() -> Self {
178 Self::new_unchecked(M::mod_one())
179 }
180}
181
182impl<M> Add for MInt<M>
183where
184 M: MIntBase,
185{
186 type Output = Self;
187 #[inline]
188 fn add(self, rhs: Self) -> Self::Output {
189 Self::new_unchecked(M::mod_add(self.x, rhs.x))
190 }
191}
192impl<M> Sub for MInt<M>
193where
194 M: MIntBase,
195{
196 type Output = Self;
197 #[inline]
198 fn sub(self, rhs: Self) -> Self::Output {
199 Self::new_unchecked(M::mod_sub(self.x, rhs.x))
200 }
201}
202impl<M> Mul for MInt<M>
203where
204 M: MIntBase,
205{
206 type Output = Self;
207 #[inline]
208 fn mul(self, rhs: Self) -> Self::Output {
209 Self::new_unchecked(M::mod_mul(self.x, rhs.x))
210 }
211}
212impl<M> Div for MInt<M>
213where
214 M: MIntBase,
215{
216 type Output = Self;
217 #[inline]
218 fn div(self, rhs: Self) -> Self::Output {
219 Self::new_unchecked(M::mod_div(self.x, rhs.x))
220 }
221}
222impl<M> Neg for MInt<M>
223where
224 M: MIntBase,
225{
226 type Output = Self;
227 #[inline]
228 fn neg(self) -> Self::Output {
229 Self::new_unchecked(M::mod_neg(self.x))
230 }
231}
232impl<M> Sum for MInt<M>
233where
234 M: MIntBase,
235{
236 #[inline]
237 fn sum<I: Iterator<Item = Self>>(iter: I) -> Self {
238 iter.fold(<Self as Zero>::zero(), Add::add)
239 }
240}
241impl<M> Product for MInt<M>
242where
243 M: MIntBase,
244{
245 #[inline]
246 fn product<I: Iterator<Item = Self>>(iter: I) -> Self {
247 iter.fold(<Self as One>::one(), Mul::mul)
248 }
249}
250impl<'a, M: 'a> Sum<&'a MInt<M>> for MInt<M>
251where
252 M: MIntBase,
253{
254 #[inline]
255 fn sum<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
256 iter.fold(<Self as Zero>::zero(), Add::add)
257 }
258}
259impl<'a, M: 'a> Product<&'a MInt<M>> for MInt<M>
260where
261 M: MIntBase,
262{
263 #[inline]
264 fn product<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
265 iter.fold(<Self as One>::one(), Mul::mul)
266 }
267}
268impl<M> Display for MInt<M>
269where
270 M: MIntBase<Inner: Display>,
271{
272 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
273 write!(f, "{}", self.inner())
274 }More examples
crates/competitive/src/num/mint/mod.rs (line 64)
63 fn fast_print<W: std::io::Write>(&self, writer: &mut FastOutput<W>) {
64 self.inner().fast_print(writer);
65 }
66}
67
68#[codesnip::entry(when("MIntBase", "coding"))]
69impl<M> SerdeByteStr for MInt<M>
70where
71 M: MIntBase<Inner: SerdeByteStr>,
72{
73 fn serialize(&self, buf: &mut Vec<u8>) {
74 self.inner().serialize(buf)
75 }crates/competitive/src/math/number_theoretic_transform.rs (line 792)
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 }Source§impl MInt<DynModuloU32>
impl MInt<DynModuloU32>
Sourcepub fn set_mod(m: u32)
pub fn set_mod(m: u32)
Examples found in repository?
crates/library_checker/src/number_theory/sqrt_mod.rs (line 9)
5pub fn sqrt_mod(reader: impl Read, writer: impl Write) {
6 prepare_io!(reader, writer);
7 sc!(q, yp: [(u32, u32); iter q]);
8 for (y, p) in yp {
9 DynMIntU32::set_mod(p);
10 if let Some(x) = DynMIntU32::from(y).sqrt() {
11 pp!(x);
12 } else {
13 pp!("-1");
14 }
15 }
16}More examples
Source§impl<M> MInt<M>
impl<M> MInt<M>
Sourcepub unsafe fn matrix_product_avx2(
a: &[Vec<Self>],
b: &[Vec<Self>],
scale: u32,
) -> Vec<Vec<Self>>
pub unsafe fn matrix_product_avx2( a: &[Vec<Self>], b: &[Vec<Self>], scale: u32, ) -> Vec<Vec<Self>>
§Safety
AVX2 must be available. The modulus must be odd, greater than one and below 2^30.
scale must be below the modulus; entries must have canonical raw representations.
Examples found in repository?
crates/competitive/src/num/mint/montgomery_dot_product.rs (line 22)
9 fn try_matrix_product(
10 _a: &[Vec<MInt<Self>>],
11 _b: &[Vec<MInt<Self>>],
12 ) -> Option<Vec<Vec<MInt<Self>>>> {
13 #[cfg(target_arch = "x86_64")]
14 if _a.len() >= 32
15 && _b.len() >= 32
16 && _b[0].len() >= 32
17 && <Self as MontgomeryReduction32>::MOD > 1
18 && <Self as MontgomeryReduction32>::MOD < 1 << 30
19 && <Self as MontgomeryReduction32>::MOD % 2 == 1
20 && is_x86_feature_detected!("avx2")
21 {
22 return Some(unsafe { MInt::matrix_product_avx2(_a, _b, 1) });
23 }
24 None
25 }Trait Implementations§
Source§impl<M> AddAssign for MInt<M>where
M: MIntBase,
impl<M> AddAssign for MInt<M>where
M: MIntBase,
Source§fn add_assign(&mut self, rhs: MInt<M>)
fn add_assign(&mut self, rhs: MInt<M>)
Performs the
+= operation. Read moreSource§impl<M> AddAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
impl<M> AddAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
Source§fn add_assign(&mut self, other: &MInt<M>)
fn add_assign(&mut self, other: &MInt<M>)
Performs the
+= operation. Read moreimpl<M> Copy for MInt<M>where
M: MIntBase,
Source§impl<M> DivAssign for MInt<M>where
M: MIntBase,
impl<M> DivAssign for MInt<M>where
M: MIntBase,
Source§fn div_assign(&mut self, rhs: MInt<M>)
fn div_assign(&mut self, rhs: MInt<M>)
Performs the
/= operation. Read moreSource§impl<M> DivAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
impl<M> DivAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
Source§fn div_assign(&mut self, other: &MInt<M>)
fn div_assign(&mut self, other: &MInt<M>)
Performs the
/= operation. Read moreSource§impl<M> DotProduct for MInt<M>where
M: MIntDotProduct,
impl<M> DotProduct for MInt<M>where
M: MIntDotProduct,
fn try_matrix_product( a: &[Vec<Self>], b: &[Vec<Self>], ) -> Option<Vec<Vec<Self>>>
fn dot_product(x: &[Self], y: &[Self]) -> Self
fn add_scaled_assign(x: &mut [Self], y: &[Self], a: &Self)
impl<M> Eq for MInt<M>where
M: MIntBase,
Source§impl<M> FastPrint for MInt<M>
impl<M> FastPrint for MInt<M>
fn fast_print<W: Write>(&self, writer: &mut FastOutput<W>)
Source§impl<M> FormalPowerSeriesCoefficient for MInt<M>
impl<M> FormalPowerSeriesCoefficient for MInt<M>
type Base = M
fn pow(self, exp: usize) -> Self
fn memorized_fact(mf: &MemorizedFactorial<Self::Base>) -> &[Self]
fn memorized_inv_fact(mf: &MemorizedFactorial<Self::Base>) -> &[Self]
fn memorized_inv(mf: &MemorizedFactorial<Self::Base>, n: usize) -> Self
fn signed_pow(self, exp: isize) -> Self
fn memorized_factorial(n: usize) -> MemorizedFactorial<Self::Base>
Source§impl<M> FormalPowerSeriesCoefficientSqrt for MInt<M>
impl<M> FormalPowerSeriesCoefficientSqrt for MInt<M>
fn sqrt_coefficient(&self) -> Option<Self>
Source§impl<M> MulAssign for MInt<M>where
M: MIntBase,
impl<M> MulAssign for MInt<M>where
M: MIntBase,
Source§fn mul_assign(&mut self, rhs: MInt<M>)
fn mul_assign(&mut self, rhs: MInt<M>)
Performs the
*= operation. Read moreSource§impl<M> MulAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
impl<M> MulAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
Source§fn mul_assign(&mut self, other: &MInt<M>)
fn mul_assign(&mut self, other: &MInt<M>)
Performs the
*= operation. Read moreSource§impl<M> RandomSpec<MInt<M>> for RangeFull
impl<M> RandomSpec<MInt<M>> for RangeFull
Source§impl<M> SerdeByteStr for MInt<M>where
M: MIntBase<Inner: SerdeByteStr>,
impl<M> SerdeByteStr for MInt<M>where
M: MIntBase<Inner: SerdeByteStr>,
Source§impl<M> SubAssign for MInt<M>where
M: MIntBase,
impl<M> SubAssign for MInt<M>where
M: MIntBase,
Source§fn sub_assign(&mut self, rhs: MInt<M>)
fn sub_assign(&mut self, rhs: MInt<M>)
Performs the
-= operation. Read moreSource§impl<M> SubAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
impl<M> SubAssign<&MInt<M>> for MInt<M>where
M: MIntBase,
Source§fn sub_assign(&mut self, other: &MInt<M>)
fn sub_assign(&mut self, other: &MInt<M>)
Performs the
-= operation. Read moreAuto Trait Implementations§
impl<M> Freeze for MInt<M>
impl<M> RefUnwindSafe for MInt<M>
impl<M> Send for MInt<M>
impl<M> Sync for MInt<M>
impl<M> Unpin for MInt<M>
impl<M> UnsafeUnpin for MInt<M>
impl<M> UnwindSafe for MInt<M>
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Mutably borrows from an owned value. Read more