1use super::{Xorshift, miller_rabin, primitive_root};
2
3const K: usize = 1 << 21;
4const POW_BLOCK_BITS: u32 = 15;
5const POW_BLOCK: usize = 1 << POW_BLOCK_BITS;
6const FRAC_SHIFT: u32 = 10;
7const FRAC_LEN: usize = 1 + (1 << (30 - FRAC_SHIFT));
8const FRAC_DEN_LIMIT: u16 = 1 << 11;
9const BSGS_SIZE: usize = 1 << 17;
10
11const fn is_direct_table_mod<const P: u32>() -> bool {
12 P as usize <= K
13}
14
15const fn table_k<const P: u32>() -> usize {
16 if is_direct_table_mod::<P>() {
17 P as usize - 1
18 } else {
19 K
20 }
21}
22
23const fn table_len<const P: u32>() -> usize {
24 2 * table_k::<P>() + 1
25}
26
27pub struct FastPrimeMod<const P: u32, const BUILD_INV: bool = true, const BUILD_POW: bool = true> {
29 root: u32,
30 pow_lo: Box<[u32]>,
31 pow_hi: Box<[u32]>,
32 frac: Box<[u32]>,
33 log: Box<[u32]>,
34 inv: Box<[u32]>,
35}
36
37impl<const P: u32, const BUILD_INV: bool, const BUILD_POW: bool>
38 FastPrimeMod<P, BUILD_INV, BUILD_POW>
39{
40 pub fn new() -> Self {
46 assert!(
47 BUILD_INV || BUILD_POW,
48 "at least one of BUILD_INV or BUILD_POW must be true"
49 );
50 assert!(P < 1 << 30, "P must be less than 2^30");
51 assert!(
52 P % 2 == 1 && miller_rabin(P as u64),
53 "P must be an odd prime"
54 );
55
56 let (root, pow_lo, pow_hi, log) = if BUILD_POW {
57 let root = if P == 998_244_353 {
58 3
59 } else {
60 primitive_root(P as u64) as u32
61 };
62 let (pow_lo, pow_hi) = build_pow::<P>(root);
63 let log = build_log::<P>(root, &pow_lo, &pow_hi);
64 (root, pow_lo, pow_hi, log)
65 } else {
66 (
67 0,
68 Vec::new().into_boxed_slice(),
69 Vec::new().into_boxed_slice(),
70 Vec::new().into_boxed_slice(),
71 )
72 };
73 let inv = if BUILD_INV {
74 build_inv::<P>()
75 } else {
76 Vec::new().into_boxed_slice()
77 };
78 let frac = if is_direct_table_mod::<P>() {
79 Vec::new().into_boxed_slice()
80 } else {
81 build_frac::<P>()
82 };
83 Self {
84 root,
85 pow_lo,
86 pow_hi,
87 frac,
88 log,
89 inv,
90 }
91 }
92
93 #[inline]
95 pub fn modulus(&self) -> u32 {
96 P
97 }
98
99 #[inline(always)]
100 fn small_fraction(&self, x: u32) -> (usize, u32) {
101 let k = table_k::<P>();
102 if is_direct_table_mod::<P>() {
103 debug_assert!(1 <= x && x < P);
104 return (k + x as usize, 1);
105 }
106 let packed = self.frac[(x >> FRAC_SHIFT) as usize];
107 let a = packed >> 16;
108 let b = packed & 0xffff;
109 let t = x.wrapping_mul(b).wrapping_sub(a.wrapping_mul(P));
110 debug_assert!({
111 let t = x as i64 * b as i64 - a as i64 * P as i64;
112 t != 0 && -(k as i64) <= t && t <= k as i64
113 });
114 ((k as u32).wrapping_add(t) as usize, b)
115 }
116}
117
118impl<const P: u32, const BUILD_POW: bool> FastPrimeMod<P, true, BUILD_POW> {
119 #[inline]
125 pub fn inverse(&self, x: u32) -> u32 {
126 assert!(1 <= x && x < P);
127 let (i, b) = self.small_fraction(x);
128 mul_mod_raw::<P>(self.inv[i], b)
129 }
130}
131
132impl<const P: u32, const BUILD_INV: bool> FastPrimeMod<P, BUILD_INV, true> {
133 #[inline]
135 pub fn primitive_root(&self) -> u32 {
136 self.root
137 }
138
139 #[inline]
147 pub fn pow(&self, a: u32, exp: u64) -> u32 {
148 assert!(a < P);
149 if a == 0 {
150 return if exp == 0 { 1 } else { 0 };
151 }
152 let ord = (P - 1) as u64;
153 self.pow_nonzero_reduced(a, (exp % ord) as u32)
154 }
155
156 #[inline]
162 pub fn pow_nonzero_reduced(&self, a: u32, exp_mod: u32) -> u32 {
163 assert!(1 <= a && a < P);
164 assert!(exp_mod < P - 1);
165 let exp = (self.log_r(a) as u64 * exp_mod as u64 % (P - 1) as u64) as u32;
166 self.pow_root_reduced(exp)
167 }
168
169 #[inline]
175 pub fn pow_root_reduced(&self, exp_mod: u32) -> u32 {
176 assert!(exp_mod < P - 1);
177 pow_root_raw::<P>(exp_mod, &self.pow_lo, &self.pow_hi)
178 }
179
180 #[inline]
181 fn log_r(&self, x: u32) -> u32 {
182 let (i, b) = self.small_fraction(x);
183 let k = table_k::<P>();
184 let ord = P - 1;
185 self.log[i] + ord - self.log[k + b as usize]
186 }
187}
188
189impl<const P: u32, const BUILD_INV: bool, const BUILD_POW: bool> Default
190 for FastPrimeMod<P, BUILD_INV, BUILD_POW>
191{
192 fn default() -> Self {
193 Self::new()
194 }
195}
196
197fn build_pow<const P: u32>(root: u32) -> (Box<[u32]>, Box<[u32]>) {
198 let mut pow_lo = vec![0; POW_BLOCK + 1].into_boxed_slice();
199 let mut pow_hi = vec![0; POW_BLOCK + 1].into_boxed_slice();
200 pow_lo[0] = 1;
201 pow_hi[0] = 1;
202 for i in 0..POW_BLOCK {
203 pow_lo[i + 1] = mul_mod_raw::<P>(pow_lo[i], root);
204 }
205 let block_power = pow_lo[POW_BLOCK];
206 for i in 0..POW_BLOCK {
207 pow_hi[i + 1] = mul_mod_raw::<P>(pow_hi[i], block_power);
208 }
209 (pow_lo, pow_hi)
210}
211
212fn build_inv<const P: u32>() -> Box<[u32]> {
213 let k = table_k::<P>();
214 let mut inv = vec![0; table_len::<P>()].into_boxed_slice();
215 inv[k + 1] = 1;
216 for i in 2..=k {
217 let q = P.div_ceil(i as u32);
218 let r = i as u32 * q - P;
219 inv[k + i] = mul_mod_raw::<P>(inv[k + r as usize], q);
220 }
221 for i in 1..=k {
222 inv[k - i] = P - inv[k + i];
223 }
224 inv
225}
226
227fn build_log<const P: u32>(root: u32, pow_lo: &[u32], pow_hi: &[u32]) -> Box<[u32]> {
228 let k = table_k::<P>();
229 let ord = P - 1;
230 let mut lpf = vec![0; k + 1].into_boxed_slice();
231 let mut primes = vec![];
232 lpf[1] = 1;
233 for i in 2..=k {
234 if lpf[i] == 0 {
235 lpf[i] = i as u32;
236 primes.push(i as u32);
237 }
238 for &p in primes.iter() {
239 let p = p as usize;
240 if p > lpf[i] as usize || p > k / i {
241 break;
242 }
243 lpf[i * p] = p as u32;
244 }
245 }
246
247 let baby_size = (BSGS_SIZE as u32).min(ord);
248 let mut baby = U32Map::new(baby_size as usize);
249 let mut pw = 1;
250 for i in 0..baby_size {
251 baby.insert(pw, i);
252 pw = mul_mod_raw::<P>(pw, root);
253 }
254 let q = pow_root_raw::<P>(ord - baby_size, pow_lo, pow_hi);
255
256 let mut log = vec![0; table_len::<P>()].into_boxed_slice();
257 log[k + 1] = 0;
258 let mut rng = Xorshift::default();
259 let small_primes = [2, 3, 5, 7, 11, 13, 17, 19];
260 for i in 2..=k {
261 let p = lpf[i] as usize;
262 if p < i {
263 log[k + i] = add_mod(log[k + p], log[k + i / p], ord);
264 } else if i < 100 {
265 let mut x = i as u32;
266 let mut ans = 0;
267 loop {
268 if let Some(v) = baby.get(x) {
269 log[k + i] = ans + v;
270 break;
271 }
272 ans += baby_size;
273 x = mul_mod_raw::<P>(x, q);
274 }
275 } else if i > P as usize / i {
276 let j = (P as usize) / i;
277 let r = (P as usize) % i;
278 let x = add_mod(log[k + r], ord / 2, ord);
279 let y = log[k + j];
280 log[k + i] = if x >= y { x - y } else { x + ord - y };
281 } else {
282 loop {
283 let exp = rng.rand(ord as u64) as u32;
284 let mut ans = if exp == 0 { 0 } else { ord - exp };
285 let mut x = mul_mod_raw::<P>(i as u32, pow_root_raw::<P>(exp, pow_lo, pow_hi));
286 for q in small_primes {
287 while x.is_multiple_of(q) {
288 x /= q;
289 ans = add_mod(ans, log[k + q as usize], ord);
290 }
291 }
292 if x as usize >= k {
293 continue;
294 }
295 while (i as u32) < x && lpf[x as usize] < i as u32 {
296 let q = lpf[x as usize];
297 x /= q;
298 ans = add_mod(ans, log[k + q as usize], ord);
299 }
300 if 1 < x && x < i as u32 {
301 ans = add_mod(ans, log[k + x as usize], ord);
302 x = 1;
303 }
304 if x == 1 {
305 log[k + i] = ans;
306 break;
307 }
308 }
309 }
310 }
311 for i in 1..=k {
312 log[k - i] = add_mod(log[k + i], ord / 2, ord);
313 }
314 log
315}
316
317fn build_frac<const P: u32>() -> Box<[u32]> {
318 let mut frac = vec![0; FRAC_LEN].into_boxed_slice();
319 let mut stack = vec![(0u16, 1u16, 1u16, 1u16)];
320 while let Some((a, b, c, d)) = stack.pop() {
321 let nb = b + d;
322 if nb < FRAC_DEN_LIMIT {
323 stack.push((a + c, nb, c, d));
324 stack.push((a, b, a + c, nb));
325 } else {
326 let s = (a as u64 * P as u64 / ((1u64 << FRAC_SHIFT) * b as u64)) as usize;
327 let t = (c as u64 * P as u64 / ((1u64 << FRAC_SHIFT) * d as u64)) as usize;
328 frac[s] = pack_frac(a, b);
329 frac[t] = pack_frac(c, d);
330 let a = a.min(c);
331 let b = b.min(d);
332 if s + 1 < t {
333 for x in &mut frac[s + 1..t] {
334 *x = pack_frac(a, b);
335 }
336 }
337 }
338 }
339 frac
340}
341
342#[inline(always)]
343fn pow_root_raw<const P: u32>(exp: u32, pow_lo: &[u32], pow_hi: &[u32]) -> u32 {
344 let lo = exp as usize & (POW_BLOCK - 1);
345 let hi = exp as usize >> POW_BLOCK_BITS;
346 mul_mod_raw::<P>(pow_lo[lo], pow_hi[hi])
347}
348
349#[inline(always)]
350fn mul_mod_raw<const P: u32>(a: u32, b: u32) -> u32 {
351 (a as u64 * b as u64 % P as u64) as u32
352}
353
354#[inline]
355fn add_mod(a: u32, b: u32, m: u32) -> u32 {
356 let c = a + b;
357 if c >= m { c - m } else { c }
358}
359
360#[inline]
361fn pack_frac(a: u16, b: u16) -> u32 {
362 (a as u32) << 16 | b as u32
363}
364
365struct U32Map {
366 keys: Box<[u32]>,
367 values: Box<[u32]>,
368 mask: usize,
369}
370
371impl U32Map {
372 fn new(capacity: usize) -> Self {
373 let len = (capacity * 2).next_power_of_two();
374 Self {
375 keys: vec![0; len].into_boxed_slice(),
376 values: vec![0; len].into_boxed_slice(),
377 mask: len - 1,
378 }
379 }
380
381 fn insert(&mut self, key: u32, value: u32) {
382 debug_assert_ne!(key, 0);
383 let mut i = self.index(key);
384 while self.keys[i] != 0 && self.keys[i] != key {
385 i = (i + 1) & self.mask;
386 }
387 self.keys[i] = key;
388 self.values[i] = value;
389 }
390
391 fn get(&self, key: u32) -> Option<u32> {
392 debug_assert_ne!(key, 0);
393 let mut i = self.index(key);
394 while self.keys[i] != 0 {
395 if self.keys[i] == key {
396 return Some(self.values[i]);
397 }
398 i = (i + 1) & self.mask;
399 }
400 None
401 }
402
403 #[inline]
404 fn index(&self, key: u32) -> usize {
405 let mut x = key as u64;
406 x = (x ^ (x >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
407 x = (x ^ (x >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
408 ((x ^ (x >> 31)) as usize) & self.mask
409 }
410}
411
412#[cfg(test)]
413mod tests {
414 use super::*;
415 use crate::{math::modinv, tools::RandomSpec};
416
417 fn mod_pow<const P: u32>(a: u32, mut exp: u64) -> u32 {
418 let mut x = a;
419 let mut y = 1;
420 while exp > 0 {
421 if exp & 1 == 1 {
422 y = mul_mod_raw::<P>(y, x);
423 }
424 x = mul_mod_raw::<P>(x, x);
425 exp >>= 1;
426 }
427 y
428 }
429
430 fn check_prime_mod<const P: u32>() {
431 let all = FastPrimeMod::<P>::new();
432 let inv_only = FastPrimeMod::<P, true, false>::new();
433 let pow_only = FastPrimeMod::<P, false, true>::new();
434 assert_eq!(all.modulus(), P);
435 assert_eq!(inv_only.modulus(), P);
436 assert_eq!(pow_only.modulus(), P);
437 assert!(1 <= all.primitive_root() && all.primitive_root() < P);
438 assert_eq!(all.primitive_root(), pow_only.primitive_root());
439
440 for x in [1, 2, 3, P / 2, P - 2, P - 1, 1_234_567] {
441 if x < P {
442 let expected = modinv(x as u64, P as u64) as u32;
443 assert_eq!(all.inverse(x), expected);
444 assert_eq!(inv_only.inverse(x), expected);
445 assert_eq!(mul_mod_raw::<P>(x, all.inverse(x)), 1);
446 }
447 }
448
449 let fixed = [
450 (0, 0),
451 (0, 1),
452 (1, u64::MAX),
453 (2, 0),
454 (2, 1),
455 (2, 123_456_789_012_345),
456 (P - 1, 0),
457 (P - 1, 1),
458 (P - 1, 2),
459 (P - 1, 123_456_789),
460 ];
461 for (a, exp) in fixed {
462 let expected = mod_pow::<P>(a, exp);
463 assert_eq!(all.pow(a, exp), expected);
464 assert_eq!(pow_only.pow(a, exp), expected);
465 if a != 0 {
466 let exp_mod = (exp % (P - 1) as u64) as u32;
467 assert_eq!(all.pow_nonzero_reduced(a, exp_mod), expected);
468 assert_eq!(pow_only.pow_nonzero_reduced(a, exp_mod), expected);
469 }
470 }
471
472 let mut rng = Xorshift::default();
473 for x in (1..P).rand_iter(&mut rng).take(2_000) {
474 let expected = modinv(x as u64, P as u64) as u32;
475 assert_eq!(all.inverse(x), expected);
476 assert_eq!(inv_only.inverse(x), expected);
477 }
478 for (a, exp) in (0..P, 0u64..).rand_iter(&mut rng).take(2_000) {
479 let expected = mod_pow::<P>(a, exp);
480 assert_eq!(all.pow(a, exp), expected);
481 assert_eq!(pow_only.pow(a, exp), expected);
482 if a != 0 {
483 let exp_mod = (exp % (P - 1) as u64) as u32;
484 assert_eq!(all.pow_nonzero_reduced(a, exp_mod), expected);
485 assert_eq!(pow_only.pow_nonzero_reduced(a, exp_mod), expected);
486 }
487 }
488 for exp_mod in (0..P - 1).rand_iter(&mut rng).take(2_000) {
489 assert_eq!(
490 pow_only.pow_root_reduced(exp_mod),
491 mod_pow::<P>(pow_only.primitive_root(), exp_mod as u64)
492 );
493 }
494 }
495
496 #[test]
497 fn test_fast_prime_mod() {
498 check_prime_mod::<998_244_353>();
499 check_prime_mod::<1_000_000_007>();
500 check_prime_mod::<3>();
501 check_prime_mod::<101>();
502 check_prime_mod::<2_017>();
503 check_prime_mod::<1_000_003>();
504 check_prime_mod::<2_097_143>();
505 check_prime_mod::<2_097_169>();
506 }
507}