Skip to main content

competitive/num/mint/
montgomery.rs

1use super::*;
2
3impl<M> MIntBase for M
4where
5    M: MontgomeryReduction32,
6{
7    type Inner = u32;
8    fn get_mod() -> Self::Inner {
9        <Self as MontgomeryReduction32>::MOD
10    }
11    fn mod_zero() -> Self::Inner {
12        0
13    }
14    fn mod_one() -> Self::Inner {
15        Self::N1
16    }
17    fn mod_add(x: Self::Inner, y: Self::Inner) -> Self::Inner {
18        let z = x + y;
19        let m = Self::get_mod();
20        if z >= m { z - m } else { z }
21    }
22    fn mod_sub(x: Self::Inner, y: Self::Inner) -> Self::Inner {
23        if x < y {
24            x + Self::get_mod() - y
25        } else {
26            x - y
27        }
28    }
29    fn mod_mul(x: Self::Inner, y: Self::Inner) -> Self::Inner {
30        Self::reduce(x as u64 * y as u64)
31    }
32
33    fn mod_div(x: Self::Inner, y: Self::Inner) -> Self::Inner {
34        Self::mod_mul(x, Self::mod_inv(y))
35    }
36    fn mod_neg(x: Self::Inner) -> Self::Inner {
37        if x == 0 { 0 } else { Self::get_mod() - x }
38    }
39    fn mod_inv(x: Self::Inner) -> Self::Inner {
40        let p = Self::get_mod() as i32;
41        let (mut a, mut b) = (x as i32, p);
42        let (mut u, mut x) = (1, 0);
43        while a != 0 {
44            let k = b / a;
45            x -= k * u;
46            b -= k * a;
47            std::mem::swap(&mut x, &mut u);
48            std::mem::swap(&mut b, &mut a);
49        }
50        Self::reduce((if x < 0 { x + p } else { x }) as u64 * Self::N3 as u64)
51    }
52    fn mod_inner(x: Self::Inner) -> Self::Inner {
53        Self::reduce(x as u64)
54    }
55}
56impl<M> MIntConvert<u32> for M
57where
58    M: MontgomeryReduction32,
59{
60    fn from(x: u32) -> Self::Inner {
61        Self::reduce(x as u64 * Self::N2 as u64)
62    }
63    fn into(x: Self::Inner) -> u32 {
64        Self::reduce(x as u64)
65    }
66    fn mod_into() -> u32 {
67        <Self as MIntBase>::get_mod()
68    }
69}
70impl<M> MIntConvert<u64> for M
71where
72    M: MontgomeryReduction32,
73{
74    fn from(x: u64) -> Self::Inner {
75        Self::reduce(x % Self::get_mod() as u64 * Self::N2 as u64)
76    }
77    fn into(x: Self::Inner) -> u64 {
78        Self::reduce(x as u64) as u64
79    }
80    fn mod_into() -> u64 {
81        <Self as MIntBase>::get_mod() as u64
82    }
83}
84impl<M> MIntConvert<usize> for M
85where
86    M: MontgomeryReduction32,
87{
88    fn from(x: usize) -> Self::Inner {
89        Self::reduce(x as u64 % Self::get_mod() as u64 * Self::N2 as u64)
90    }
91    fn into(x: Self::Inner) -> usize {
92        Self::reduce(x as u64) as usize
93    }
94    fn mod_into() -> usize {
95        <Self as MIntBase>::get_mod() as usize
96    }
97}
98impl<M> MIntConvert<i32> for M
99where
100    M: MontgomeryReduction32,
101{
102    fn from(x: i32) -> Self::Inner {
103        let x = x % <Self as MIntBase>::get_mod() as i32;
104        let x = if x < 0 {
105            (x + <Self as MIntBase>::get_mod() as i32) as u64
106        } else {
107            x as u64
108        };
109        Self::reduce(x * Self::N2 as u64)
110    }
111    fn into(x: Self::Inner) -> i32 {
112        Self::reduce(x as u64) as i32
113    }
114    fn mod_into() -> i32 {
115        <Self as MIntBase>::get_mod() as i32
116    }
117}
118impl<M> MIntConvert<i64> for M
119where
120    M: MontgomeryReduction32,
121{
122    fn from(x: i64) -> Self::Inner {
123        let x = x % <Self as MIntBase>::get_mod() as i64;
124        let x = if x < 0 {
125            (x + <Self as MIntBase>::get_mod() as i64) as u64
126        } else {
127            x as u64
128        };
129        Self::reduce(x * Self::N2 as u64)
130    }
131    fn into(x: Self::Inner) -> i64 {
132        Self::reduce(x as u64) as i64
133    }
134    fn mod_into() -> i64 {
135        <Self as MIntBase>::get_mod() as i64
136    }
137}
138impl<M> MIntConvert<isize> for M
139where
140    M: MontgomeryReduction32,
141{
142    fn from(x: isize) -> Self::Inner {
143        let x = x % <Self as MIntBase>::get_mod() as isize;
144        let x = if x < 0 {
145            (x + <Self as MIntBase>::get_mod() as isize) as u64
146        } else {
147            x as u64
148        };
149        Self::reduce(x * Self::N2 as u64)
150    }
151    fn into(x: Self::Inner) -> isize {
152        Self::reduce(x as u64) as isize
153    }
154    fn mod_into() -> isize {
155        <Self as MIntBase>::get_mod() as isize
156    }
157}
158/// m is prime, n = 2^32
159pub trait MontgomeryReduction32 {
160    /// m
161    const MOD: u32;
162    /// (-m)^{-1} mod n
163    const R: u32 = {
164        let m = Self::MOD;
165        let mut r = 0;
166        let mut t = 0;
167        let mut i = 0;
168        while i < 32 {
169            if t % 2 == 0 {
170                t += m;
171                r += 1 << i;
172            }
173            t /= 2;
174            i += 1;
175        }
176        r
177    };
178    /// n^1 mod m
179    const N1: u32 = ((1u64 << 32) % Self::MOD as u64) as _;
180    /// n^2 mod m
181    const N2: u32 = (Self::N1 as u64 * Self::N1 as u64 % Self::MOD as u64) as _;
182    /// n^3 mod m
183    const N3: u32 = (Self::N1 as u64 * Self::N2 as u64 % Self::MOD as u64) as _;
184    /// n^{-1}x = (x + (xr mod n)m) / n
185    fn reduce(x: u64) -> u32 {
186        let m: u32 = Self::MOD;
187        let r = Self::R;
188        let mut x = ((x + r.wrapping_mul(x as u32) as u64 * m as u64) >> 32) as u32;
189        if x >= m {
190            x -= m;
191        }
192        x
193    }
194}
195macro_rules! define_montgomery_reduction_32 {
196    ($([$name:ident, $m:expr, $mint_name:ident $(,)?]),* $(,)?) => {
197        $(
198            pub enum $name {}
199            impl MontgomeryReduction32 for $name {
200                const MOD: u32 = $m;
201            }
202            pub type $mint_name = MInt<$name>;
203        )*
204    };
205}
206define_montgomery_reduction_32!(
207    [Modulo167772161, 167_772_161, MInt167772161],
208    [Modulo469762049, 469_762_049, MInt469762049],
209    [Modulo754974721, 754_974_721, MInt754974721],
210    [Modulo998244353, 998_244_353, MInt998244353],
211);
212
213#[cfg(test)]
214mod tests {
215    use super::*;
216    use crate::num::montgomery::MInt998244353 as M;
217    use crate::tools::Xorshift;
218
219    #[test]
220    fn test_mint998244353() {
221        let mut rng = Xorshift::default();
222        const Q: usize = 1000;
223        assert_eq!(0, MInt998244353::zero().inner());
224        assert_eq!(1, MInt998244353::one().inner());
225        assert_eq!(
226            Modulo998244353::reduce(Modulo998244353::N3 as u64),
227            Modulo998244353::N2
228        );
229        assert_eq!(
230            Modulo998244353::reduce(Modulo998244353::N2 as u64),
231            Modulo998244353::N1
232        );
233        assert_eq!(Modulo998244353::reduce(Modulo998244353::N1 as u64), 1);
234        for _ in 0..Q {
235            let x = rng.random(..MInt998244353::get_mod());
236            assert_eq!(x, MInt998244353::new(x).inner());
237            assert_eq!((-M::new(x)).inner(), (-MInt998244353::new(x)).inner());
238            assert_eq!(x, MInt998244353::new(x).inv().inv().inner());
239            assert_eq!(M::new(x).inv().inner(), MInt998244353::new(x).inv().inner());
240        }
241
242        for _ in 0..Q {
243            let x = rng.random(..MInt998244353::get_mod());
244            let y = rng.random(..MInt998244353::get_mod());
245            assert_eq!(
246                (M::new(x) + M::new(y)).inner(),
247                (MInt998244353::new(x) + MInt998244353::new(y)).inner()
248            );
249            assert_eq!(
250                (M::new(x) - M::new(y)).inner(),
251                (MInt998244353::new(x) - MInt998244353::new(y)).inner()
252            );
253            assert_eq!(
254                (M::new(x) * M::new(y)).inner(),
255                (MInt998244353::new(x) * MInt998244353::new(y)).inner()
256            );
257            assert_eq!(
258                (M::new(x) / M::new(y)).inner(),
259                (MInt998244353::new(x) / MInt998244353::new(y)).inner()
260            );
261            assert_eq!(
262                M::new(x).pow(y as usize).inner(),
263                MInt998244353::new(x).pow(y as usize).inner()
264            );
265        }
266
267        for _ in 0..Q {
268            let x = rng.rand64();
269            assert_eq!(
270                M::from(x as u32).inner(),
271                MInt998244353::from(x as u32).inner()
272            );
273            assert_eq!(M::from(x).inner(), MInt998244353::from(x).inner());
274            assert_eq!(
275                M::from(x as usize).inner(),
276                MInt998244353::from(x as usize).inner()
277            );
278            assert_eq!(
279                M::from(x as i32).inner(),
280                MInt998244353::from(x as i32).inner()
281            );
282            assert_eq!(
283                M::from(x as i64).inner(),
284                MInt998244353::from(x as i64).inner()
285            );
286            assert_eq!(
287                M::from(x as isize).inner(),
288                MInt998244353::from(x as isize).inner()
289            );
290        }
291    }
292}