1use super::*;
2
3use std::{
4 fmt::{self, Debug, Display},
5 hash::{Hash, Hasher},
6 iter::{Product, Sum},
7 marker::PhantomData,
8 ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign},
9 str::FromStr,
10};
11
12#[repr(transparent)]
13pub struct MInt<M>
14where
15 M: MIntBase,
16{
17 x: M::Inner,
18 _marker: PhantomData<fn() -> M>,
19}
20
21pub trait MIntConvert<T = <Self as MIntBase>::Inner>: MIntBase {
22 fn from(x: T) -> <Self as MIntBase>::Inner;
23 fn into(x: <Self as MIntBase>::Inner) -> T;
24 fn mod_into() -> T;
25}
26
27pub trait MIntBase {
28 type Inner: Sized + Copy + Eq + Debug + Hash;
29 fn get_mod() -> Self::Inner;
30 fn mod_zero() -> Self::Inner;
31 fn mod_one() -> Self::Inner;
32 fn mod_add(x: Self::Inner, y: Self::Inner) -> Self::Inner;
33 fn mod_sub(x: Self::Inner, y: Self::Inner) -> Self::Inner;
34 fn mod_mul(x: Self::Inner, y: Self::Inner) -> Self::Inner;
35 fn mod_div(x: Self::Inner, y: Self::Inner) -> Self::Inner;
36 fn mod_neg(x: Self::Inner) -> Self::Inner;
37 fn mod_inv(x: Self::Inner) -> Self::Inner;
38 fn mod_pow(x: Self::Inner, y: usize) -> Self::Inner {
39 let (mut x, mut y, mut z) = (x, y, Self::mod_one());
40 while y > 0 {
41 if y & 1 == 1 {
42 z = Self::mod_mul(z, x);
43 }
44 x = Self::mod_mul(x, x);
45 y >>= 1;
46 }
47 z
48 }
49 fn mod_inner(x: Self::Inner) -> Self::Inner {
50 x
51 }
52}
53
54impl<M> MInt<M>
55where
56 M: MIntConvert,
57{
58 #[inline]
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 }
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 }
275}
276impl<M> FromStr for MInt<M>
277where
278 M: MIntConvert + MIntBase<Inner: FromStr>,
279{
280 type Err = <M::Inner as FromStr>::Err;
281 #[inline]
282 fn from_str(s: &str) -> Result<Self, Self::Err> {
283 s.parse::<M::Inner>().map(Self::new)
284 }
285}
286impl<M> Scan for MInt<M>
287where
288 M: MIntConvert + MIntBase<Inner: Scan<Output = M::Inner>>,
289{
290 type Output = Self;
291 #[inline]
292 fn scan<S: ScanSource>(source: &mut S) -> Option<Self> {
293 M::Inner::scan(source).map(Self::new)
294 }
295}
296macro_rules! impl_mint_ref_binop {
297 ($imp:ident, $method:ident, $t:ty) => {
298 impl<M> $imp<$t> for &$t
299 where
300 M: MIntBase,
301 {
302 type Output = <$t as $imp<$t>>::Output;
303 #[inline]
304 fn $method(self, other: $t) -> <$t as $imp<$t>>::Output {
305 $imp::$method(*self, other)
306 }
307 }
308 impl<M> $imp<&$t> for $t
309 where
310 M: MIntBase,
311 {
312 type Output = <$t as $imp<$t>>::Output;
313 #[inline]
314 fn $method(self, other: &$t) -> <$t as $imp<$t>>::Output {
315 $imp::$method(self, *other)
316 }
317 }
318 impl<M> $imp<&$t> for &$t
319 where
320 M: MIntBase,
321 {
322 type Output = <$t as $imp<$t>>::Output;
323 #[inline]
324 fn $method(self, other: &$t) -> <$t as $imp<$t>>::Output {
325 $imp::$method(*self, *other)
326 }
327 }
328 };
329}
330impl_mint_ref_binop!(Add, add, MInt<M>);
331impl_mint_ref_binop!(Sub, sub, MInt<M>);
332impl_mint_ref_binop!(Mul, mul, MInt<M>);
333impl_mint_ref_binop!(Div, div, MInt<M>);
334macro_rules! impl_mint_ref_unop {
335 ($imp:ident, $method:ident, $t:ty) => {
336 impl<M> $imp for &$t
337 where
338 M: MIntBase,
339 {
340 type Output = <$t as $imp>::Output;
341 #[inline]
342 fn $method(self) -> <$t as $imp>::Output {
343 $imp::$method(*self)
344 }
345 }
346 };
347}
348impl_mint_ref_unop!(Neg, neg, MInt<M>);
349macro_rules! impl_mint_ref_op_assign {
350 ($imp:ident, $method:ident, $t:ty, $fromimp:ident, $frommethod:ident) => {
351 impl<M> $imp<$t> for $t
352 where
353 M: MIntBase,
354 {
355 #[inline]
356 fn $method(&mut self, rhs: $t) {
357 *self = $fromimp::$frommethod(*self, rhs);
358 }
359 }
360 impl<M> $imp<&$t> for $t
361 where
362 M: MIntBase,
363 {
364 #[inline]
365 fn $method(&mut self, other: &$t) {
366 $imp::$method(self, *other);
367 }
368 }
369 };
370}
371impl_mint_ref_op_assign!(AddAssign, add_assign, MInt<M>, Add, add);
372impl_mint_ref_op_assign!(SubAssign, sub_assign, MInt<M>, Sub, sub);
373impl_mint_ref_op_assign!(MulAssign, mul_assign, MInt<M>, Mul, mul);
374impl_mint_ref_op_assign!(DivAssign, div_assign, MInt<M>, Div, div);
375
376#[cfg(test)]
377mod tests {
378 use super::{MInt, MIntBase, MIntConvert};
379 #[cfg(target_arch = "x86_64")]
380 use crate::tools::avx512_enabled;
381 use crate::{
382 algebra::{AddMulOperation, DotProduct},
383 define_basic_mint32, define_basic_mintbase, impl_basic_mint_dot_product,
384 math::Matrix,
385 num::{MIntDotProduct, mint_basic, montgomery},
386 tools::Xorshift,
387 };
388 use std::mem::swap;
389
390 define_basic_mint32!([Modulo17, 17, MInt17]);
391 impl_basic_mint_dot_product!(u32, u64; Modulo17);
392
393 #[test]
394 fn test_random_matrix_products() {
395 let mut rng = Xorshift::new_with_seed(374938);
396 macro_rules! check {
397 ($mint:ty) => {{
398 let (n, m, p) = (rng.random(1..160), rng.random(1..160), rng.random(1..160));
399 let a: Vec<Vec<$mint>> = (0..n)
400 .map(|_| (0..m).map(|_| rng.random(..)).collect())
401 .collect();
402 let b: Vec<Vec<$mint>> = (0..m)
403 .map(|_| (0..p).map(|_| rng.random(..)).collect())
404 .collect();
405 let a: Matrix<AddMulOperation<_>> = Matrix::from_vec(a);
406 let b = Matrix::from_vec(b);
407 let result = &a * &b;
408 for i in 0..n {
409 for j in 0..p {
410 let expected = (0..m).map(|k| a[i][k] * b[k][j]).sum();
411 assert_eq!(result[i][j], expected);
412 }
413 }
414 }};
415 }
416 for _ in 0..24 {
417 check!(MInt17);
418 check!(mint_basic::MInt998244353);
419 check!(mint_basic::MInt1000000007);
420 mint_basic::DynMIntU32::set_mod(rng.random(1u32..1 << 29) * 2 + 1);
421 check!(mint_basic::DynMIntU32);
422 check!(montgomery::MInt167772161);
423 check!(montgomery::MInt469762049);
424 check!(montgomery::MInt754974721);
425 check!(montgomery::MInt998244353);
426 }
427 mint_basic::DynMIntU32::set_mod(1_000_000_007);
428 }
429
430 #[test]
431 fn test_random_vector_operations() {
432 let mut rng = Xorshift::new_with_seed(629714);
433 macro_rules! check {
434 ($mint:ty) => {{
435 let n = rng.random(0..2048);
436 let mut x: Vec<$mint> = (0..n).map(|_| rng.random(..)).collect();
437 let y: Vec<$mint> = (0..n).map(|_| rng.random(..)).collect();
438 let a: $mint = rng.random(..);
439 let expected: Vec<_> = x.iter().zip(&y).map(|(&x, &y)| x + a * y).collect();
440 let dot = x.iter().zip(&y).map(|(&x, &y)| x * y).sum();
441 assert_eq!(<$mint>::dot_product(&x, &y), dot);
442 <$mint>::add_scaled_assign(&mut x, &y, &a);
443 assert_eq!(x, expected);
444 }};
445 }
446 for _ in 0..256 {
447 check!(MInt17);
448 check!(mint_basic::MInt2);
449 check!(mint_basic::MInt998244353);
450 check!(mint_basic::MInt1000000007);
451 mint_basic::DynMIntU32::set_mod(rng.random(1..));
452 check!(mint_basic::DynMIntU32);
453 mint_basic::DynMIntU64::set_mod(rng.random(1..));
454 check!(mint_basic::DynMIntU64);
455 check!(montgomery::MInt167772161);
456 check!(montgomery::MInt469762049);
457 check!(montgomery::MInt754974721);
458 check!(montgomery::MInt998244353);
459 }
460 mint_basic::DynMIntU32::set_mod(1_000_000_007);
461 mint_basic::DynMIntU64::set_mod(1_000_000_007);
462 }
463}