1use super::*;
2use std::{cell::UnsafeCell, mem::swap};
3
4#[macro_export]
5macro_rules! define_basic_mintbase {
6 ($name:ident, $m:expr, $basety:ty, $signedty:ty, $upperty:ty, [$($unsigned:ty),*], [$($signed:ty),*]) => {
7 pub enum $name {}
8 impl MIntBase for $name {
9 type Inner = $basety;
10 #[inline]
11 fn get_mod() -> Self::Inner {
12 $m
13 }
14 #[inline]
15 fn mod_zero() -> Self::Inner {
16 0
17 }
18 #[inline]
19 fn mod_one() -> Self::Inner {
20 (Self::get_mod() != 1) as $basety
21 }
22 #[inline]
23 fn mod_add(x: Self::Inner, y: Self::Inner) -> Self::Inner {
24 let m = Self::get_mod();
25 let (z, borrow) = x.overflowing_sub(m - y);
26 if borrow { z.wrapping_add(m) } else { z }
27 }
28 #[inline]
29 fn mod_sub(x: Self::Inner, y: Self::Inner) -> Self::Inner {
30 if x < y {
31 Self::get_mod() - (y - x)
32 } else {
33 x - y
34 }
35 }
36 #[inline]
37 fn mod_mul(x: Self::Inner, y: Self::Inner) -> Self::Inner {
38 $name::rem(x as $upperty * y as $upperty) as $basety
40 }
41
42
43
44 #[inline]
45 fn mod_div(x: Self::Inner, y: Self::Inner) -> Self::Inner {
46 Self::mod_mul(x, Self::mod_inv(y))
47 }
48 #[inline]
49 fn mod_neg(x: Self::Inner) -> Self::Inner {
50 if x == 0 {
51 0
52 } else {
53 Self::get_mod() - x
54 }
55 }
56 fn mod_inv(x: Self::Inner) -> Self::Inner {
57 let (mut a, mut b) = (x, Self::get_mod());
58 let (mut u, mut v) = (1, 0);
59 let mut negative = true;
60 while a != 0 {
62 let k = b / a;
63 v += k * u;
64 b -= k * a;
65 swap(&mut u, &mut v);
66 swap(&mut b, &mut a);
67 negative = !negative;
68 }
69 if negative { Self::mod_neg(v) } else { v }
70 }
71 }
72
73 $(impl MIntConvert<$unsigned> for $name {
74 #[inline]
75 fn from(x: $unsigned) -> Self::Inner {
76 (x % <Self as MIntBase>::get_mod() as $unsigned) as $basety
77 }
78 #[inline]
79 fn into(x: Self::Inner) -> $unsigned {
80 x as $unsigned
81 }
82 #[inline]
83 fn mod_into() -> $unsigned {
84 <Self as MIntBase>::get_mod() as $unsigned
85 }
86 })*
87 $(impl MIntConvert<$signed> for $name {
88 #[inline]
89 fn from(x: $signed) -> Self::Inner {
90 let modulus = (<Self as MIntBase>::get_mod() as $signed).cast_unsigned();
91 let value = (x.unsigned_abs() % modulus) as $basety;
92 if x < 0 {
93 <Self as MIntBase>::mod_neg(value)
94 } else {
95 value
96 }
97 }
98 #[inline]
99 fn into(x: Self::Inner) -> $signed {
100 x as $signed
101 }
102 #[inline]
103 fn mod_into() -> $signed {
104 <Self as MIntBase>::get_mod() as $signed
105 }
106 })*
107 };
108}
109
110#[macro_export]
111macro_rules! define_basic_mint32 {
112 ($([$name:ident, $m:expr, $mint_name:ident]),*) => {
113 $(define_basic_mintbase!(
114 $name,
115 $m,
116 u32,
117 i32,
118 u64,
119 [u32, u64, u128, usize],
120 [i32, i64, i128, isize]
121 );
122 impl $name {
123 fn rem(x: u64) -> u64 {
124 x % $m
125 }
126 }
127 pub type $mint_name = MInt<$name>;)*
128 };
129}
130
131thread_local!(static DYN_MODULUS_U32: UnsafeCell<BarrettReduction<u64>> = const { UnsafeCell::new(BarrettReduction::<u64>::new_with_im(1_000_000_007, !0 / 1_000_000_007)) });
132impl DynModuloU32 {
133 pub fn set_mod(m: u32) {
134 DYN_MODULUS_U32
135 .with(|cell| unsafe { *cell.get() = BarrettReduction::<u64>::new(m as u64) });
136 }
137 fn rem(x: u64) -> u64 {
138 DYN_MODULUS_U32.with(|cell| unsafe { (*cell.get()).rem(x) })
139 }
140}
141impl DynMIntU32 {
142 pub fn set_mod(m: u32) {
143 DynModuloU32::set_mod(m)
144 }
145}
146
147thread_local!(static DYN_MODULUS_U64: UnsafeCell<BarrettReduction<u128>> = const { UnsafeCell::new(BarrettReduction::<u128>::new_with_im(1_000_000_007, !0 / 1_000_000_007)) });
148impl DynModuloU64 {
149 pub fn set_mod(m: u64) {
150 DYN_MODULUS_U64
151 .with(|cell| unsafe { *cell.get() = BarrettReduction::<u128>::new(m as u128) })
152 }
153 fn rem(x: u128) -> u128 {
154 DYN_MODULUS_U64.with(|cell| unsafe { (*cell.get()).rem(x) })
155 }
156}
157impl DynMIntU64 {
158 pub fn set_mod(m: u64) {
159 DynModuloU64::set_mod(m)
160 }
161}
162
163define_basic_mint32!(
164 [Modulo998244353, 998_244_353, MInt998244353],
165 [Modulo1000000007, 1_000_000_007, MInt1000000007],
166 [Modulo1000000009, 1_000_000_009, MInt1000000009]
167);
168
169define_basic_mintbase!(
170 DynModuloU32,
171 DYN_MODULUS_U32.with(|cell| unsafe { (*cell.get()).get_mod() as u32 }),
172 u32,
173 i32,
174 u64,
175 [u32, u64, u128, usize],
176 [i32, i64, i128, isize]
177);
178pub type DynMIntU32 = MInt<DynModuloU32>;
179define_basic_mintbase!(
180 DynModuloU64,
181 DYN_MODULUS_U64.with(|cell| unsafe { (*cell.get()).get_mod() as u64 }),
182 u64,
183 i64,
184 u128,
185 [u64, u128, usize],
186 [i64, i128, isize]
187);
188pub type DynMIntU64 = MInt<DynModuloU64>;
189
190pub struct Modulo2;
191impl MIntBase for Modulo2 {
192 type Inner = u32;
193 #[inline]
194 fn get_mod() -> Self::Inner {
195 2
196 }
197 #[inline]
198 fn mod_zero() -> Self::Inner {
199 0
200 }
201 #[inline]
202 fn mod_one() -> Self::Inner {
203 1
204 }
205 #[inline]
206 fn mod_add(x: Self::Inner, y: Self::Inner) -> Self::Inner {
207 x ^ y
208 }
209 #[inline]
210 fn mod_sub(x: Self::Inner, y: Self::Inner) -> Self::Inner {
211 x ^ y
212 }
213 #[inline]
214 fn mod_mul(x: Self::Inner, y: Self::Inner) -> Self::Inner {
215 x & y
216 }
217 #[inline]
218 fn mod_div(x: Self::Inner, y: Self::Inner) -> Self::Inner {
219 assert_ne!(y, 0);
220 x
221 }
222 #[inline]
223 fn mod_neg(x: Self::Inner) -> Self::Inner {
224 x
225 }
226 #[inline]
227 fn mod_inv(x: Self::Inner) -> Self::Inner {
228 assert_ne!(x, 0);
229 x
230 }
231 #[inline]
232 fn mod_pow(x: Self::Inner, y: usize) -> Self::Inner {
233 if y == 0 { 1 } else { x }
234 }
235}
236macro_rules! impl_to_mint_base_for_modulo2 {
237 ($name:ident, $basety:ty, [$($t:ty),*]) => {
238 $(impl MIntConvert<$t> for $name {
239 #[inline]
240 fn from(x: $t) -> Self::Inner {
241 (x & 1) as $basety
242 }
243 #[inline]
244 fn into(x: Self::Inner) -> $t {
245 x as $t
246 }
247 #[inline]
248 fn mod_into() -> $t {
249 2
250 }
251 })*
252 };
253}
254impl_to_mint_base_for_modulo2!(
255 Modulo2,
256 u32,
257 [
258 u8, u16, u32, u64, u128, usize, i8, i16, i32, i64, i128, isize
259 ]
260);
261pub type MInt2 = MInt<Modulo2>;
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266 use crate::tools::Xorshift;
267
268 macro_rules! test_mint {
269 ($test_name:ident $mint:ident $($m:expr)?) => {
270 #[test]
271 fn $test_name() {
272 let mut rng = Xorshift::default();
273 const Q: usize = 10_000;
274 for _ in 0..Q {
275 $($mint::set_mod(rng.gen(..$m));)?
276 let a = $mint::new_unchecked(rng.random(1..$mint::get_mod()));
277 let x = a.inv();
278 assert!(x.inner() < $mint::get_mod());
279 assert_eq!(a * x, $mint::one());
280 }
281 for _ in 0..100 {
282 let n = rng.random(0..100);
283 let x: Vec<$mint> = (0..n).map(|_| rng.random(..)).collect();
284 let y: Vec<$mint> = (0..n).map(|_| rng.random(..)).collect();
285 assert_eq!(
286 $mint::dot_product(&x, &y),
287 x.iter().zip(&y).map(|(&x, &y)| x * y).sum()
288 );
289 }
290 for n in 0..=576 {
291 let x = vec![$mint::new_unchecked($mint::get_mod() - 1); n];
292 assert_eq!(
293 $mint::dot_product(&x, &x),
294 x.iter().map(|&x| x * x).sum()
295 );
296 }
297 }
298 };
299 }
300 test_mint!(test_mint2 MInt2);
301 test_mint!(test_mint998244353 MInt998244353);
302 test_mint!(test_mint1000000007 MInt1000000007);
303 test_mint!(test_mint1000000009 MInt1000000009);
304
305 macro_rules! test_dyn_mint_arithmetic {
306 ($name:ident, $mint:ident, $int:ty, $signed:ty) => {
307 #[test]
308 fn $name() {
309 let mut rng = Xorshift::default();
310 let moduli: Vec<_> = [1, 2, <$int>::MAX]
311 .into_iter()
312 .chain(rng.random_iter(1..).take(100))
313 .collect();
314 for modulus in moduli {
315 $mint::set_mod(modulus);
316 assert_eq!($mint::one().inner() as u128, 1 % modulus as u128);
317 let values: Vec<_> = [0, modulus - 1]
318 .into_iter()
319 .chain(rng.random_iter(0..modulus).take(16))
320 .collect();
321 for &x in &values {
322 if modulus > 1 && crate::math::gcd(x as u64, modulus as u64) == 1 {
323 let inverse = $mint::from(x).inv().inner();
324 assert!(inverse < modulus);
325 assert_eq!(x as u128 * inverse as u128 % modulus as u128, 1);
326 }
327 for &y in &values {
328 let a = $mint::from(x);
329 let b = $mint::from(y);
330 let modulus = modulus as u128;
331 assert_eq!((a + b).inner() as u128, (x as u128 + y as u128) % modulus);
332 assert_eq!(
333 (a - b).inner() as u128,
334 (x as u128 + modulus - y as u128) % modulus
335 );
336 }
337 }
338 for x in [<$signed>::MIN, -1, 0, 1, <$signed>::MAX]
339 .into_iter()
340 .chain(rng.random_iter(..).take(16))
341 {
342 assert_eq!(
343 $mint::from(x).inner() as i128,
344 (x as i128).rem_euclid(modulus as i128)
345 );
346 }
347 for x in [i128::MIN, -1, 0, 1, i128::MAX]
348 .into_iter()
349 .chain(rng.random_iter(..).take(16))
350 {
351 assert_eq!(
352 $mint::from(x).inner() as i128,
353 x.rem_euclid(modulus as i128)
354 );
355 }
356 let lengths: Vec<_> = [0, 1, 63, 64, 65, 511, 512, 513, 600]
357 .into_iter()
358 .chain(rng.random_iter(0..600).take(8))
359 .collect();
360 for n in lengths {
361 let x: Vec<$mint> = rng.random_iter(..).take(n).collect();
362 let y: Vec<$mint> = rng.random_iter(..).take(n).collect();
363 let expected = x.iter().zip(&y).fold(0, |sum, (x, y)| {
364 (sum + x.inner() as u128 * y.inner() as u128) % modulus as u128
365 });
366 assert_eq!($mint::dot_product(&x, &y).inner() as u128, expected);
367 }
368 }
369 $mint::set_mod(1_000_000_007);
370 }
371 };
372 }
373 test_dyn_mint_arithmetic!(test_dyn_mint_u32_arithmetic, DynMIntU32, u32, i32);
374 test_dyn_mint_arithmetic!(test_dyn_mint_u64_arithmetic, DynMIntU64, u64, i64);
375
376 #[test]
377 fn test_dyn_mint_u32_dot_product() {
378 DynMIntU32::set_mod(1_000_000_007);
379 let mut rng = Xorshift::default();
380 for n in 0..=576 {
381 let x: Vec<DynMIntU32> = rng.random_iter(..).take(n).collect();
382 let y: Vec<DynMIntU32> = rng.random_iter(..).take(n).collect();
383 assert_eq!(
384 DynMIntU32::dot_product(&x, &y),
385 x.iter().zip(&y).map(|(&x, &y)| x * y).sum()
386 );
387 }
388 DynMIntU32::set_mod(u32::MAX);
389 for n in 0..=576 {
390 let x = vec![DynMIntU32::new_unchecked(u32::MAX - 1); n];
391 assert_eq!(
392 DynMIntU32::dot_product(&x, &x),
393 x.iter().map(|&x| x * x).sum()
394 );
395 }
396 DynMIntU32::set_mod(1_000_000_007);
397 }
398}