competitive/num/mint/
montgomery.rs1use 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}
158pub trait MontgomeryReduction32 {
160 const MOD: u32;
162 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 const N1: u32 = ((1u64 << 32) % Self::MOD as u64) as _;
180 const N2: u32 = (Self::N1 as u64 * Self::N1 as u64 % Self::MOD as u64) as _;
182 const N3: u32 = (Self::N1 as u64 * Self::N2 as u64 % Self::MOD as u64) as _;
184 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}