Skip to main content

competitive/num/mint/
mint_basic.rs

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                // (x as $upperty * y as $upperty % Self::get_mod() as $upperty) as $basety
39                $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                // b * u + a * v equals the modulus, so coefficient updates fit.
61                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}