1use super::{
2 ConvolveSteps, MInt, MIntBase, MIntConvert, One, Zero, advise_huge_pages,
3 fast_fourier_transform::ConvolveRealFft, montgomery::*,
4};
5#[cfg(target_arch = "x86_64")]
6use super::{
7 SimdBackend,
8 mint_fft_convolve::{convolve_mint_avx2, convolve_u64_avx2},
9 montgomery_simd, simd_backend,
10};
11use std::{
12 cell::UnsafeCell,
13 marker::PhantomData,
14 num::Wrapping,
15 ops::{AddAssign, Mul, SubAssign},
16};
17
18#[cfg(target_arch = "x86_64")]
19#[inline]
20fn batch_ntt_simd_backend(len: usize, width: usize) -> SimdBackend {
21 if width < 8 || !width.is_power_of_two() || len < width * 4 {
23 SimdBackend::Scalar
24 } else if width == 8 {
25 if is_x86_feature_detected!("avx2") {
26 SimdBackend::Avx2
27 } else {
28 SimdBackend::Scalar
29 }
30 } else {
31 simd_backend()
32 }
33}
34
35pub struct Convolve<M>(PhantomData<fn() -> M>);
36pub type Convolve998244353 = Convolve<Modulo998244353>;
37pub type MIntConvolve<M> = Convolve<(M, (Modulo167772161, Modulo469762049, Modulo754974721))>;
40pub type U64Convolve = Convolve<(u64, (Modulo167772161, Modulo469762049, Modulo754974721))>;
43
44macro_rules! impl_ntt_modulus {
45 ($([$name:ident, $g:expr]),*) => {
46 $(
47 impl Montgomery32NttModulus for $name {}
48 )*
49 };
50}
51impl_ntt_modulus!(
52 [Modulo167772161, 3],
53 [Modulo469762049, 3],
54 [Modulo754974721, 11],
55 [Modulo998244353, 3]
56);
57
58const fn reduce(z: u64, p: u32, r: u32) -> u32 {
59 let mut z = ((z + r.wrapping_mul(z as u32) as u64 * p as u64) >> 32) as u32;
60 if z >= p {
61 z -= p;
62 }
63 z
64}
65const fn mod_mul(x: u32, y: u32, p: u32, r: u32) -> u32 {
66 reduce(x as u64 * y as u64, p, r)
67}
68const fn mod_pow(mut x: u32, mut y: u32, p: u32, r: u32, mut z: u32) -> u32 {
69 while y > 0 {
70 if y & 1 == 1 {
71 z = mod_mul(z, x, p, r);
72 }
73 x = mod_mul(x, x, p, r);
74 y >>= 1;
75 }
76 z
77}
78
79pub trait Montgomery32NttModulus: Sized + MontgomeryReduction32 {
80 const PRIMITIVE_ROOT: u32 = {
81 let mut g = 3u32;
82 loop {
83 let mut ok = true;
84 let mut d = 1u32;
85 while d * d < Self::MOD {
86 if (Self::MOD - 1) % d == 0 {
87 let ds = [d, (Self::MOD - 1) / d];
88 let mut i = 0;
89 while i < 2 {
90 ok &= ds[i] == Self::MOD - 1
91 || mod_pow(
92 reduce(g as u64 * Self::N2 as u64, Self::MOD, Self::R),
93 ds[i],
94 Self::MOD,
95 Self::R,
96 Self::N1,
97 ) != Self::N1;
98 i += 1;
99 }
100 }
101 d += 1;
102 }
103 if ok {
104 break;
105 }
106 g += 2;
107 }
108 g
109 };
110 const RANK: u32 = (Self::MOD - 1).trailing_zeros();
111 const INFO: NttInfo = NttInfo::new::<Self>();
112}
113
114#[derive(Debug, PartialEq)]
115pub struct NttInfo {
116 root: [u32; 32],
117 inv_root: [u32; 32],
118 rate3: [u32; 32],
119 rate3_packed: [[u32; 8]; 32],
120 inv_rate3_packed: [[u32; 8]; 32],
121}
122impl NttInfo {
123 const fn new<M>() -> Self
124 where
125 M: Montgomery32NttModulus,
126 {
127 let mut root = [0; 32];
128 let mut inv_root = [0; 32];
129 let mut rate3_values = [0; 32];
130 let mut rate3_packed = [[0; 8]; 32];
131 let mut inv_rate3_packed = [[0; 8]; 32];
132 let rank = M::RANK as usize;
133
134 let g = reduce(M::PRIMITIVE_ROOT as u64 * M::N2 as u64, M::MOD, M::R);
135 root[rank] = mod_pow(g, (M::MOD - 1) >> rank, M::MOD, M::R, M::N1);
136 inv_root[rank] = mod_pow(root[rank], M::MOD - 2, M::MOD, M::R, M::N1);
137 let mut i = rank - 1;
138 loop {
139 root[i] = mod_mul(root[i + 1], root[i + 1], M::MOD, M::R);
140 inv_root[i] = mod_mul(inv_root[i + 1], inv_root[i + 1], M::MOD, M::R);
141 if i == 0 {
142 break;
143 }
144 i -= 1;
145 }
146
147 let (mut i, mut prod, mut inv_prod) = (0, M::N1, M::N1);
148 while i < rank - 2 {
149 let rate3 = mod_mul(root[i + 3], prod, M::MOD, M::R);
150 rate3_values[i] = rate3;
151 let rate3_2 = mod_mul(rate3, rate3, M::MOD, M::R);
152 let rate3_3 = mod_mul(rate3_2, rate3, M::MOD, M::R);
153 let inv_rate3 = mod_mul(inv_root[i + 3], inv_prod, M::MOD, M::R);
154 let inv_rate3_2 = mod_mul(inv_rate3, inv_rate3, M::MOD, M::R);
155 let inv_rate3_3 = mod_mul(inv_rate3_2, inv_rate3, M::MOD, M::R);
156 rate3_packed[i] = [
157 rate3.wrapping_mul(M::R),
158 rate3,
159 rate3_2.wrapping_mul(M::R),
160 rate3_2,
161 rate3_3.wrapping_mul(M::R),
162 rate3_3,
163 0,
164 0,
165 ];
166 inv_rate3_packed[i] = [
167 inv_rate3.wrapping_mul(M::R),
168 inv_rate3,
169 inv_rate3_2.wrapping_mul(M::R),
170 inv_rate3_2,
171 inv_rate3_3.wrapping_mul(M::R),
172 inv_rate3_3,
173 0,
174 0,
175 ];
176 prod = mod_mul(prod, inv_root[i + 3], M::MOD, M::R);
177 inv_prod = mod_mul(inv_prod, root[i + 3], M::MOD, M::R);
178 i += 1;
179 }
180
181 NttInfo {
182 root,
183 inv_root,
184 rate3: rate3_values,
185 rate3_packed,
186 inv_rate3_packed,
187 }
188 }
189}
190
191const LAZY_THRESHOLD: u32 = 1 << 30;
192
193#[inline]
194fn add_scalar<M>(x: u32, y: u32) -> u32
195where
196 M: Montgomery32NttModulus,
197{
198 let modulus = if M::MOD < LAZY_THRESHOLD {
199 M::MOD * 2
200 } else {
201 M::MOD
202 };
203 let sum = x + y;
204 if sum >= modulus { sum - modulus } else { sum }
205}
206
207#[inline]
208fn sub_scalar<M>(x: u32, y: u32) -> u32
209where
210 M: Montgomery32NttModulus,
211{
212 let modulus = if M::MOD < LAZY_THRESHOLD {
213 M::MOD * 2
214 } else {
215 M::MOD
216 };
217 if x < y { x + modulus - y } else { x - y }
218}
219
220#[inline]
221fn mul_scalar<M>(x: u32, y: u32) -> u32
222where
223 M: Montgomery32NttModulus,
224{
225 if M::MOD < LAZY_THRESHOLD {
226 let z = x as u64 * y as u64;
227 ((z + M::R.wrapping_mul(z as u32) as u64 * M::MOD as u64) >> 32) as u32
228 } else {
229 M::mod_mul(x, y)
230 }
231}
232
233fn ntt_scalar<M>(a: &mut [MInt<M>])
234where
235 M: Montgomery32NttModulus,
236{
237 ntt_batch_scalar(a, 1);
238}
239
240fn ntt_batch<M>(a: &mut [MInt<M>], width: usize)
241where
242 M: Montgomery32NttModulus,
243{
244 #[cfg(target_arch = "x86_64")]
245 {
246 match batch_ntt_simd_backend(a.len(), width) {
247 SimdBackend::Avx512 => {
248 unsafe { ntt_simd::ntt_batch_avx512(a, width) };
250 return;
251 }
252 SimdBackend::Avx2 => {
253 unsafe { ntt_simd::ntt_batch_avx2::<_, false>(a, width) };
255 return;
256 }
257 SimdBackend::Scalar => {}
258 }
259 }
260 ntt_batch_scalar(a, width);
261}
262
263fn ntt_batch_scalar<M>(a: &mut [MInt<M>], width: usize)
264where
265 M: Montgomery32NttModulus,
266{
267 let n = a.len() / width;
268 if n <= 1 {
269 return;
270 }
271 let mut v = n / 2;
272 if n.trailing_zeros() & 1 == 1 {
273 let (l, r) = a.split_at_mut(v * width);
274 for (x0, x1) in l.iter_mut().zip(r) {
275 let a0 = *x0;
276 let a1 = *x1;
277 *x0 = a0 + a1;
278 *x1 = a0 - a1;
279 }
280 v >>= 1;
281 }
282 let imag = MInt::<M>::new_unchecked(M::INFO.root[2]);
283 while v > 1 {
284 let mut w1 = MInt::<M>::one();
285 let mut w2 = w1;
286 let mut w3 = w1;
287 for (s, a) in a.chunks_exact_mut((v << 1) * width).enumerate() {
288 let (l, r) = a.split_at_mut(v * width);
289 let (ll, lr) = l.split_at_mut((v >> 1) * width);
290 let (rl, rr) = r.split_at_mut((v >> 1) * width);
291 for (((x0, x1), x2), x3) in ll.iter_mut().zip(lr).zip(rl).zip(rr) {
292 let a0 = *x0;
293 let a1 = *x1 * w1;
294 let a2 = *x2 * w2;
295 let a3 = *x3 * w3;
296 let a0pa2 = a0 + a2;
297 let a0na2 = a0 - a2;
298 let a1pa3 = a1 + a3;
299 let a1na3imag = (a1 - a3) * imag;
300 *x0 = a0pa2 + a1pa3;
301 *x1 = a0pa2 - a1pa3;
302 *x2 = a0na2 + a1na3imag;
303 *x3 = a0na2 - a1na3imag;
304 }
305 let rate = &M::INFO.rate3_packed[s.trailing_ones() as usize];
306 w1 *= MInt::<M>::new_unchecked(rate[1]);
307 w2 *= MInt::<M>::new_unchecked(rate[3]);
308 w3 *= MInt::<M>::new_unchecked(rate[5]);
309 }
310 v >>= 2;
311 }
312}
313
314fn intt_scalar<M>(a: &mut [MInt<M>])
315where
316 M: Montgomery32NttModulus,
317{
318 intt_batch_scalar(a, 1);
319}
320
321fn intt_batch<M>(a: &mut [MInt<M>], width: usize)
322where
323 M: Montgomery32NttModulus,
324{
325 #[cfg(target_arch = "x86_64")]
326 {
327 match batch_ntt_simd_backend(a.len(), width) {
328 SimdBackend::Avx512 => {
329 unsafe { ntt_simd::intt_batch_avx512::<_, false>(a, width) };
331 return;
332 }
333 SimdBackend::Avx2 => {
334 unsafe { ntt_simd::intt_batch_avx2::<_, false>(a, width) };
336 return;
337 }
338 SimdBackend::Scalar => {}
339 }
340 }
341 intt_batch_scalar(a, width);
342}
343
344fn intt_batch_scalar<M>(a: &mut [MInt<M>], width: usize)
345where
346 M: Montgomery32NttModulus,
347{
348 let n = a.len() / width;
349 if n <= 1 {
350 return;
351 }
352 let a = unsafe { std::slice::from_raw_parts_mut(a.as_mut_ptr().cast::<u32>(), a.len()) };
355 let mut v = 1;
356 let limit = if n.trailing_zeros() & 1 == 1 {
357 n / 2
358 } else {
359 n
360 };
361 let iimag = M::INFO.inv_root[2];
362 while v < limit {
363 let mut w1 = M::N1;
364 let mut w2 = w1;
365 let mut w3 = w1;
366 for (s, a) in a.chunks_exact_mut((v << 2) * width).enumerate() {
367 let (l, r) = a.split_at_mut((v << 1) * width);
368 let (ll, lr) = l.split_at_mut(v * width);
369 let (rl, rr) = r.split_at_mut(v * width);
370 for (((x0, x1), x2), x3) in ll.iter_mut().zip(lr).zip(rl).zip(rr) {
371 let a0 = *x0;
372 let a1 = *x1;
373 let a2 = *x2;
374 let a3 = *x3;
375 let a0pa1 = add_scalar::<M>(a0, a1);
376 let a0na1 = sub_scalar::<M>(a0, a1);
377 let a2pa3 = add_scalar::<M>(a2, a3);
378 let a2na3iimag = mul_scalar::<M>(sub_scalar::<M>(a2, a3), iimag);
379 *x0 = add_scalar::<M>(a0pa1, a2pa3);
380 *x1 = mul_scalar::<M>(add_scalar::<M>(a0na1, a2na3iimag), w1);
381 *x2 = mul_scalar::<M>(sub_scalar::<M>(a0pa1, a2pa3), w2);
382 *x3 = mul_scalar::<M>(sub_scalar::<M>(a0na1, a2na3iimag), w3);
383 }
384 let rate = &M::INFO.inv_rate3_packed[s.trailing_ones() as usize];
385 w1 = M::mod_mul(w1, rate[1]);
386 w2 = M::mod_mul(w2, rate[3]);
387 w3 = M::mod_mul(w3, rate[5]);
388 }
389 v <<= 2;
390 }
391 if n.trailing_zeros() & 1 == 1 {
392 let (l, r) = a.split_at_mut(n / 2 * width);
393 for (x0, x1) in l.iter_mut().zip(r) {
394 let a0 = *x0;
395 let a1 = *x1;
396 *x0 = add_scalar::<M>(a0, a1);
397 *x1 = sub_scalar::<M>(a0, a1);
398 }
399 }
400 let inv = M::mod_inv(<M as MIntConvert<u32>>::from(n as u32));
401 for a in a {
402 *a = M::mod_mul(*a, inv);
403 }
404}
405
406fn ntt<M>(a: &mut [MInt<M>])
407where
408 M: Montgomery32NttModulus,
409{
410 #[cfg(target_arch = "x86_64")]
411 match simd_backend() {
412 SimdBackend::Avx512 => unsafe { ntt_simd::ntt_batch_avx512(a, 1) },
413 SimdBackend::Avx2 => unsafe { ntt_simd::ntt_batch_avx2::<_, true>(a, 1) },
414 SimdBackend::Scalar => ntt_scalar(a),
415 }
416 #[cfg(not(target_arch = "x86_64"))]
417 ntt_scalar(a);
418}
419
420fn intt<M>(a: &mut [MInt<M>])
421where
422 M: Montgomery32NttModulus,
423{
424 #[cfg(target_arch = "x86_64")]
425 match simd_backend() {
426 SimdBackend::Avx512 => unsafe { ntt_simd::intt_batch_avx512::<_, true>(a, 1) },
427 SimdBackend::Avx2 => unsafe { ntt_simd::intt_batch_avx2::<_, true>(a, 1) },
428 SimdBackend::Scalar => intt_scalar(a),
429 }
430 #[cfg(not(target_arch = "x86_64"))]
431 intt_scalar(a);
432}
433
434fn ntt_rows<M>(a: &mut [MInt<M>], width: usize)
435where
436 M: Montgomery32NttModulus,
437{
438 for row in a.chunks_exact_mut(width) {
439 ntt(row);
440 }
441}
442
443fn intt_rows<M>(a: &mut [MInt<M>], width: usize)
444where
445 M: Montgomery32NttModulus,
446{
447 for row in a.chunks_exact_mut(width) {
448 intt(row);
449 }
450}
451
452#[cfg(target_arch = "x86_64")]
453fn use_block_ntt<M>(len: usize) -> bool
454where
455 M: Montgomery32NttModulus,
456{
457 len >= 64 && M::MOD < LAZY_THRESHOLD && is_x86_feature_detected!("avx2")
458}
459
460fn pointwise_multiply<M>(f: &mut [MInt<M>], g: &[MInt<M>])
461where
462 M: Montgomery32NttModulus,
463{
464 assert!(f.len() <= g.len());
465 crate::avx_helper!(
466 @dispatch simd_backend, SimdBackend;
467 unsafe { ntt_simd::pointwise_multiply_avx512(f, g) },
468 unsafe { ntt_simd::pointwise_multiply_avx2(f, g) },
469 {
470 for (f, g) in f.iter_mut().zip(g.iter()) {
471 *f *= *g;
472 }
473 }
474 )
475}
476
477fn pointwise_multiply_add<M>(sum: &mut [MInt<M>], f: &[MInt<M>], g: &[MInt<M>])
478where
479 M: Montgomery32NttModulus,
480{
481 crate::avx_helper!(
482 @dispatch simd_backend, SimdBackend;
483 unsafe { ntt_simd::pointwise_multiply_add_avx512(sum, f, g) },
484 unsafe { ntt_simd::pointwise_multiply_add_avx2(sum, f, g) },
485 {
486 for ((sum, f), g) in sum.iter_mut().zip(f.iter()).zip(g.iter()) {
487 *sum += *f * *g;
488 }
489 }
490 )
491}
492
493#[cfg(target_arch = "x86_64")]
494#[allow(unsafe_op_in_unsafe_fn)] mod ntt_simd;
496
497fn convolve_naive<T>(a: &[T], b: &[T]) -> Vec<T>
498where
499 T: Copy + Zero + AddAssign<T> + Mul<Output = T>,
500{
501 if a.is_empty() && b.is_empty() {
502 return Vec::new();
503 }
504 let len = a.len() + b.len() - 1;
505 let mut c = vec![T::zero(); len];
506 let (a, b) = if a.len() >= b.len() { (a, b) } else { (b, a) };
507 for (block, a) in a.chunks(1024).enumerate() {
508 for (i, &b) in b.iter().enumerate() {
509 let start = block * 1024 + i;
510 for (a, c) in a.iter().zip(&mut c[start..start + a.len()]) {
511 *c += *a * b;
512 }
513 }
514 }
515 c
516}
517
518fn convolve_karatsuba<T>(a: &[T], b: &[T]) -> Vec<T>
519where
520 T: Copy + Zero + AddAssign<T> + SubAssign<T> + Mul<Output = T>,
521{
522 if a.len().min(b.len()) <= 30 {
523 return convolve_naive(a, b);
524 }
525 let block_len = a.len().min(b.len()).next_power_of_two();
526 if a.len().max(b.len()) > block_len * 4 {
527 let (a, b) = if a.len() >= b.len() { (a, b) } else { (b, a) };
528 let mut result = vec![T::zero(); a.len() + b.len() - 1];
529 for (i, a) in a.chunks(block_len).enumerate() {
530 for (value, product) in result[i * block_len..]
531 .iter_mut()
532 .zip(convolve_karatsuba(a, b))
533 {
534 *value += product;
535 }
536 }
537 return result;
538 }
539 let m = a.len().max(b.len()).div_ceil(2);
540 let (a0, a1) = if a.len() <= m {
541 (a, &[][..])
542 } else {
543 a.split_at(m)
544 };
545 let (b0, b1) = if b.len() <= m {
546 (b, &[][..])
547 } else {
548 b.split_at(m)
549 };
550 let f00 = convolve_karatsuba(a0, b0);
551 let f11 = convolve_karatsuba(a1, b1);
552 let mut a0a1 = a0.to_vec();
553 for (a0a1, &a1) in a0a1.iter_mut().zip(a1) {
554 *a0a1 += a1;
555 }
556 let mut b0b1 = b0.to_vec();
557 for (b0b1, &b1) in b0b1.iter_mut().zip(b1) {
558 *b0b1 += b1;
559 }
560 let mut f01 = convolve_karatsuba(&a0a1, &b0b1);
561 for (f01, &f00) in f01.iter_mut().zip(&f00) {
562 *f01 -= f00;
563 }
564 for (f01, &f11) in f01.iter_mut().zip(&f11) {
565 *f01 -= f11;
566 }
567 let mut c = vec![T::zero(); a.len() + b.len() - 1];
568 for (c, &f00) in c.iter_mut().zip(&f00) {
569 *c += f00;
570 }
571 for (c, &f01) in c[m..].iter_mut().zip(&f01) {
572 *c += f01;
573 }
574 for (c, &f11) in c[m << 1..].iter_mut().zip(&f11) {
575 *c += f11;
576 }
577 c
578}
579
580#[cold]
581fn convolve_large_ntt<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
582where
583 M: Montgomery32NttModulus,
584{
585 let len = a.len() + b.len() - 1;
586 let ntt_len = 1usize << M::RANK;
587 let block_len = ntt_len / 2;
588 let same = a == b;
589 let transform = |a: &[MInt<M>]| {
590 let mut f = Vec::with_capacity(ntt_len);
591 advise_huge_pages(&mut f);
592 f.extend_from_slice(a);
593 Convolve::<M>::transform_ntt(f, ntt_len)
594 };
595 let fa: Vec<_> = a.chunks(block_len).map(transform).collect();
596 let fb: Option<Vec<_>> = if same {
597 None
598 } else {
599 Some(b.chunks(block_len).map(transform).collect())
600 };
601 let b_blocks = fb.as_ref().map_or(fa.len(), Vec::len);
602 let mut result = vec![MInt::<M>::zero(); len];
603 for diagonal in 0..fa.len() + b_blocks - 1 {
604 let mut spectrum = vec![MInt::<M>::zero(); ntt_len];
605 let start = diagonal.saturating_sub(b_blocks - 1);
606 for i in start..=diagonal.min(fa.len() - 1) {
607 let j = diagonal - i;
608 let g = if let Some(fb) = &fb { &fb[j] } else { &fa[j] };
609 pointwise_multiply_add(&mut spectrum, &fa[i], g);
610 }
611 spectrum = Convolve::<M>::inverse_transform_ntt(spectrum, ntt_len);
612 let offset = diagonal * block_len;
613 for (result, value) in result[offset..].iter_mut().zip(spectrum) {
614 *result += value;
615 }
616 }
617 result
618}
619
620impl<M> ConvolveSteps for Convolve<M>
621where
622 M: Montgomery32NttModulus,
623{
624 const CYCLIC: bool = true;
625
626 type T = Vec<MInt<M>>;
627 type F = Vec<MInt<M>>;
628 fn length(t: &Self::T) -> usize {
629 t.len()
630 }
631 fn transform(mut t: Self::T, len: usize) -> Self::F {
632 t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
633 #[cfg(target_arch = "x86_64")]
634 if use_block_ntt::<M>(t.len()) {
635 unsafe { ntt_simd::transform_blocks_avx2(&mut t) };
636 return t;
637 }
638 ntt(&mut t);
639 t
640 }
641 fn inverse_transform(mut f: Self::F, len: usize) -> Self::T {
642 #[cfg(target_arch = "x86_64")]
643 if use_block_ntt::<M>(f.len()) {
644 unsafe { ntt_simd::inverse_transform_blocks_avx2(&mut f) };
645 f.truncate(len);
646 return f;
647 }
648 intt(&mut f);
649 f.truncate(len);
650 f
651 }
652 fn multiply(f: &mut Self::F, g: &Self::F) {
653 assert_eq!(f.len(), g.len());
654 #[cfg(target_arch = "x86_64")]
655 if use_block_ntt::<M>(f.len()) {
656 unsafe { ntt_simd::multiply_blocks_avx2(f, g) };
657 return;
658 }
659 pointwise_multiply(f, g);
660 }
661 fn square(t: Self::T, len: usize) -> Self::T {
662 let mut f = Self::transform(t, len);
663 let g = f.clone();
664 Self::multiply(&mut f, &g);
665 Self::inverse_transform(f, len)
666 }
667 fn convolve(mut a: Self::T, mut b: Self::T) -> Self::T {
668 let (threshold, naive_threshold) = (100, 60);
669 #[cfg(target_arch = "x86_64")]
670 let (threshold, naive_threshold) = if use_block_ntt::<M>(64) {
671 (
672 60,
673 if M::RANK >= 13 && a.len().max(b.len()) <= 4096 {
674 18
675 } else if M::RANK >= 19 && a.len().max(b.len()) <= 262144 {
676 32
677 } else {
678 34
679 },
680 )
681 } else {
682 (threshold, naive_threshold)
683 };
684 if Self::length(&a).max(Self::length(&b)) <= threshold {
685 return convolve_karatsuba(&a, &b);
686 }
687 if Self::length(&a).min(Self::length(&b)) <= naive_threshold {
688 return convolve_naive(&a, &b);
689 }
690 let len = (Self::length(&a) + Self::length(&b)).saturating_sub(1);
691 let size = len.max(1).next_power_of_two();
692 let max_size = 1usize << M::RANK;
693 #[cfg(target_arch = "x86_64")]
694 let max_size = if use_block_ntt::<M>(size) {
695 max_size << 3
696 } else {
697 max_size
698 };
699 if size > max_size {
700 return convolve_large_ntt(a, b);
701 }
702 if len <= size / 2 + 2 {
703 let xa = a.pop().unwrap();
704 let xb = b.pop().unwrap();
705 let mut c = vec![MInt::<M>::zero(); len];
706 *c.last_mut().unwrap() = xa * xb;
707 for (a, c) in a.iter().zip(&mut c[b.len()..]) {
708 *c += *a * xb;
709 }
710 for (b, c) in b.iter().zip(&mut c[a.len()..]) {
711 *c += *b * xa;
712 }
713 let d = Self::convolve(a, b);
714 for (d, c) in d.into_iter().zip(&mut c) {
715 *c += d;
716 }
717 return c;
718 }
719 let same = a == b;
720 #[cfg(target_arch = "x86_64")]
721 if use_block_ntt::<M>(size) {
722 a.reserve(size - a.len());
723 b.reserve(size - b.len());
724 advise_huge_pages(&mut a);
725 advise_huge_pages(&mut b);
726 a.resize_with(size, Zero::zero);
727 b.resize_with(size, Zero::zero);
728 unsafe { ntt_simd::convolve_blocks_avx2(&mut a, &mut b, same) };
729 a.truncate(len);
730 return a;
731 }
732 let mut a = Self::transform(a, len);
733 if same {
734 for a in a.iter_mut() {
735 *a *= *a;
736 }
737 } else {
738 let b = Self::transform(b, len);
739 Self::multiply(&mut a, &b);
740 }
741 Self::inverse_transform(a, len)
742 }
743}
744
745type MVec<M> = Vec<MInt<M>>;
746
747fn convert_crt_input<M, N1, N2, N3>(t: MVec<M>, capacity: usize) -> (MVec<N1>, MVec<N2>, MVec<N3>)
748where
749 M: MIntConvert<u32>,
750 N1: Montgomery32NttModulus,
751 N2: Montgomery32NttModulus,
752 N3: Montgomery32NttModulus,
753{
754 let mut f = (
755 MVec::<N1>::with_capacity(capacity),
756 MVec::<N2>::with_capacity(capacity),
757 MVec::<N3>::with_capacity(capacity),
758 );
759 advise_huge_pages(&mut f.0);
760 advise_huge_pages(&mut f.1);
761 advise_huge_pages(&mut f.2);
762 for t in t {
763 let t: u32 = t.into();
764 f.0.push(t.into());
765 f.1.push(t.into());
766 f.2.push(t.into());
767 }
768 f
769}
770
771fn reconstruct_mint_crt<M, N1, N2, N3>(f: (MVec<N1>, MVec<N2>, MVec<N3>)) -> MVec<M>
772where
773 M: MIntConvert + MIntConvert<u32>,
774 N1: Montgomery32NttModulus,
775 N2: Montgomery32NttModulus,
776 N3: Montgomery32NttModulus,
777{
778 let t1 = MInt::<N2>::new(N1::get_mod()).inv();
779 let m1_3 = MInt::<N3>::new(N1::get_mod());
780 let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
781 let modulus = <M as MIntConvert<u32>>::mod_into() as u64;
782 let m1 = N1::get_mod() as u64;
783 let m2 = m1 * N2::get_mod() as u64 % modulus;
784 let fits_u64 = (N1::get_mod() - 1) as u128
785 + (N2::get_mod() - 1) as u128 * m1 as u128
786 + (N3::get_mod() - 1) as u128 * m2 as u128
787 <= u64::MAX as u128;
788 f.0.into_iter()
789 .zip(f.1)
790 .zip(f.2)
791 .map(|((c1, c2), c3)| {
792 let d1 = c1.inner();
793 let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
794 let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
795 let d3 = ((c3 - x) * t2).inner();
796 let value = if fits_u64 {
797 (d1 as u64 + d2 as u64 * m1 + d3 as u64 * m2) % modulus
798 } else {
799 ((d1 as u128 + d2 as u128 * m1 as u128 + d3 as u128 * m2 as u128) % modulus as u128)
800 as u64
801 };
802 MInt::<M>::from(value as u32)
803 })
804 .collect()
805}
806
807impl<M, N1, N2, N3> ConvolveSteps for Convolve<(M, (N1, N2, N3))>
808where
809 M: MIntConvert + MIntConvert<u32>,
810 N1: Montgomery32NttModulus,
811 N2: Montgomery32NttModulus,
812 N3: Montgomery32NttModulus,
813{
814 type T = MVec<M>;
815 type F = (MVec<N1>, MVec<N2>, MVec<N3>);
816 fn length(t: &Self::T) -> usize {
817 t.len()
818 }
819 fn transform(t: Self::T, len: usize) -> Self::F {
820 let npot = len.max(1).next_power_of_two();
821 let f = convert_crt_input(t, npot);
822 (
823 Convolve::<N1>::transform(f.0, npot),
824 Convolve::<N2>::transform(f.1, npot),
825 Convolve::<N3>::transform(f.2, npot),
826 )
827 }
828 fn inverse_transform(f: Self::F, len: usize) -> Self::T {
829 reconstruct_mint_crt((
830 Convolve::<N1>::inverse_transform(f.0, len),
831 Convolve::<N2>::inverse_transform(f.1, len),
832 Convolve::<N3>::inverse_transform(f.2, len),
833 ))
834 }
835 fn multiply(f: &mut Self::F, g: &Self::F) {
836 Convolve::<N1>::multiply(&mut f.0, &g.0);
837 Convolve::<N2>::multiply(&mut f.1, &g.1);
838 Convolve::<N3>::multiply(&mut f.2, &g.2);
839 }
840 fn convolve(a: Self::T, b: Self::T) -> Self::T {
841 let max_len = Self::length(&a).max(Self::length(&b));
842 let min_len = Self::length(&a).min(Self::length(&b));
843 let (balanced, short) = crate::avx_helper!(@dispatch_avx2_fma (30, 10), (384, 128));
844 if max_len <= balanced || min_len <= short {
845 return convolve_karatsuba(&a, &b);
846 }
847 let fft_limit = crate::avx_helper!(@dispatch_avx2_fma
849 1usize << ((1u64 << 50) / <M as MIntConvert<u32>>::mod_into() as u64).ilog2().min(20), 0);
850 let convolve = |a: Self::T, b: Self::T| {
851 let fft_len = (a.len() + b.len() - 1).next_power_of_two();
852 if fft_len <= 256 && a.len() * b.len() <= fft_len * 8 {
853 return convolve_karatsuba(&a, &b);
854 }
855 if fft_len <= fft_limit {
856 crate::avx_helper!(@dispatch_avx2_fma return unsafe {
857 convolve_mint_avx2(a, b)
858 }, ());
859 }
860 convolve_mint_crt::<M, N1, N2, N3>(a, b)
861 };
862 let block_len = min_len.next_power_of_two() * 8 - min_len + 1;
863 let block_len = if min_len <= fft_limit / 2 {
864 block_len.min(fft_limit - min_len + 1)
865 } else {
866 block_len
867 };
868 if max_len <= block_len {
869 return convolve(a, b);
870 }
871 let (a, b) = if a.len() >= b.len() { (a, b) } else { (b, a) };
872 let mut result = vec![MInt::<M>::zero(); a.len() + b.len() - 1];
873 for (i, a) in a.chunks(block_len).enumerate() {
874 let product = convolve(a.to_vec(), b.clone());
875 for (value, product) in result[i * block_len..].iter_mut().zip(product) {
876 *value += product;
877 }
878 }
879 result
880 }
881}
882
883fn convolve_mint_crt<M, N1, N2, N3>(a: MVec<M>, b: MVec<M>) -> MVec<M>
884where
885 M: MIntConvert + MIntConvert<u32>,
886 N1: Montgomery32NttModulus,
887 N2: Montgomery32NttModulus,
888 N3: Montgomery32NttModulus,
889{
890 let convolve = |a: MVec<M>, b: MVec<M>| {
891 let a_len = a.len();
892 let b_len = b.len();
893 let a = convert_crt_input(a, a_len);
894 let b = convert_crt_input(b, b_len);
895 reconstruct_mint_crt((
896 Convolve::<N1>::convolve(a.0, b.0),
897 Convolve::<N2>::convolve(a.1, b.1),
898 Convolve::<N3>::convolve(a.2, b.2),
899 ))
900 };
901 let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
902 let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
903 if a.len().min(b.len()) as u128 * (modulus - 1).pow(2) < capacity {
904 return convolve(a, b);
905 }
906 let block_len = ((capacity - 1) / (modulus - 1).pow(2)) as usize;
907 if block_len == 0 {
908 return convolve_naive(&a, &b);
909 }
910 let mut result = vec![MInt::<M>::zero(); a.len() + b.len() - 1];
911 for (i, a) in a.chunks(block_len).enumerate() {
912 for (j, b) in b.chunks(block_len).enumerate() {
913 let product = convolve(a.to_vec(), b.to_vec());
914 for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
915 *value += product;
916 }
917 }
918 }
919 result
920}
921
922impl<N1, N2, N3> ConvolveSteps for Convolve<(u64, (N1, N2, N3))>
923where
924 N1: Montgomery32NttModulus,
925 N2: Montgomery32NttModulus,
926 N3: Montgomery32NttModulus,
927{
928 type T = Vec<u64>;
929 type F = ([MVec<N1>; 3], [MVec<N2>; 3], [MVec<N3>; 3]);
930
931 fn length(t: &Self::T) -> usize {
932 t.len()
933 }
934
935 fn transform(t: Self::T, len: usize) -> Self::F {
936 let npot = len.max(1).next_power_of_two();
937 assert!(npot <= 1usize << N1::RANK.min(N2::RANK).min(N3::RANK));
938 assert!(
940 3 * npot as u128 * ((1u128 << 22) - 1).pow(2)
941 < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
942 );
943 let bits = if 2 * npot as u128 * (u32::MAX as u128).pow(2)
944 < N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128
945 {
946 32
947 } else {
948 22
949 };
950 let parts = if bits == 32 && t.iter().all(|&value| value <= u32::MAX as u64) {
951 1
952 } else {
953 64usize.div_ceil(bits)
954 };
955 fn split<M: Montgomery32NttModulus>(
956 t: &[u64],
957 len: usize,
958 bits: usize,
959 parts: usize,
960 ) -> [MVec<M>; 3] {
961 std::array::from_fn(|part| {
962 if part >= parts {
963 return Vec::new();
964 }
965 Convolve::<M>::transform(
966 t.iter()
967 .map(|&t| MInt::from((t >> (part * bits)) & ((1u64 << bits) - 1)))
968 .collect(),
969 len,
970 )
971 })
972 }
973 (
974 split(&t, npot, bits, parts),
975 split(&t, npot, bits, parts),
976 split(&t, npot, bits, parts),
977 )
978 }
979
980 fn inverse_transform(f: Self::F, len: usize) -> Self::T {
981 let bits = if f.0[2].is_empty() { 32 } else { 22 };
982 let t1 = MInt::<N2>::new(N1::get_mod()).inv();
983 let m1 = N1::get_mod() as u64;
984 let m1_3 = MInt::<N3>::new(N1::get_mod());
985 let t2 = (m1_3 * MInt::<N3>::new(N2::get_mod())).inv();
986 let m2 = m1 * N2::get_mod() as u64;
987 let mut result = vec![0u64; len.min(f.0[0].len())];
988 for (part, ((f1, f2), f3)) in f.0.into_iter().zip(f.1).zip(f.2).enumerate() {
989 if f1.is_empty() {
990 continue;
991 }
992 for (value, ((c1, c2), c3)) in result.iter_mut().zip(
993 Convolve::<N1>::inverse_transform(f1, len)
994 .into_iter()
995 .zip(Convolve::<N2>::inverse_transform(f2, len))
996 .zip(Convolve::<N3>::inverse_transform(f3, len)),
997 ) {
998 let d1 = c1.inner();
999 let d2 = ((c2 - MInt::<N2>::from(d1)) * t1).inner();
1000 let x = MInt::<N3>::new(d1) + MInt::<N3>::new(d2) * m1_3;
1001 let d3 = ((c3 - x) * t2).inner();
1002 let limb = (d1 as u64)
1003 .wrapping_add((d2 as u64).wrapping_mul(m1))
1004 .wrapping_add((d3 as u64).wrapping_mul(m2));
1005 *value = value.wrapping_add(limb << (part * bits));
1006 }
1007 }
1008 result
1009 }
1010
1011 fn multiply(f: &mut Self::F, g: &Self::F) {
1012 fn multiply<M: Montgomery32NttModulus>(f: &mut [MVec<M>; 3], g: &[MVec<M>; 3]) {
1013 assert_eq!(f[0].len(), g[0].len());
1014 if f[1].is_empty() || g[1].is_empty() {
1015 if f[1].is_empty() && !g[1].is_empty() {
1016 f[1] = f[0].clone();
1017 Convolve::<M>::multiply(&mut f[1], &g[1]);
1018 } else if !f[1].is_empty() {
1019 Convolve::<M>::multiply(&mut f[1], &g[0]);
1020 }
1021 Convolve::<M>::multiply(&mut f[0], &g[0]);
1022 return;
1023 }
1024 #[cfg(target_arch = "x86_64")]
1025 if use_block_ntt::<M>(f[0].len()) {
1026 for part in (1..if f[2].is_empty() { 2 } else { 3 }).rev() {
1027 let mut sum = f[0].clone();
1028 Convolve::<M>::multiply(&mut sum, &g[part]);
1029 for left in 1..=part {
1030 let mut product = f[left].clone();
1031 Convolve::<M>::multiply(&mut product, &g[part - left]);
1032 for (value, product) in sum.iter_mut().zip(product) {
1033 *value = MInt::new(value.inner() + product.inner());
1035 }
1036 }
1037 f[part] = sum;
1038 }
1039 Convolve::<M>::multiply(&mut f[0], &g[0]);
1040 return;
1041 }
1042 if f[2].is_empty() {
1043 for i in 0..f[0].len() {
1044 f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1045 f[0][i] *= g[0][i];
1046 }
1047 return;
1048 }
1049 for i in 0..f[0].len() {
1050 f[2][i] = f[0][i] * g[2][i] + f[1][i] * g[1][i] + f[2][i] * g[0][i];
1051 f[1][i] = f[0][i] * g[1][i] + f[1][i] * g[0][i];
1052 f[0][i] *= g[0][i];
1053 }
1054 }
1055 multiply(&mut f.0, &g.0);
1056 multiply(&mut f.1, &g.1);
1057 multiply(&mut f.2, &g.2);
1058 }
1059
1060 fn square(t: Self::T, len: usize) -> Self::T {
1061 let mut f = Self::transform(t, len);
1062 let g = f.clone();
1063 Self::multiply(&mut f, &g);
1064 Self::inverse_transform(f, len)
1065 }
1066
1067 fn convolve(a: Self::T, b: Self::T) -> Self::T {
1068 let max_len = Self::length(&a).max(Self::length(&b));
1069 let min_len = Self::length(&a).min(Self::length(&b));
1070 let (balanced, short) = crate::avx_helper!(@dispatch_avx2_fma (300, 64), (1536, 512));
1071 if max_len <= balanced || min_len <= short {
1072 let a_wrapping: &[Wrapping<u64>] =
1073 unsafe { std::slice::from_raw_parts(a.as_ptr().cast(), a.len()) };
1074 let b_wrapping: &[Wrapping<u64>] =
1075 unsafe { std::slice::from_raw_parts(b.as_ptr().cast(), b.len()) };
1076 let mut c = std::mem::ManuallyDrop::new(if max_len <= 300 || min_len > 60 {
1077 convolve_karatsuba(a_wrapping, b_wrapping)
1078 } else {
1079 convolve_naive(a_wrapping, b_wrapping)
1080 });
1081 return unsafe { Vec::from_raw_parts(c.as_mut_ptr().cast(), c.len(), c.capacity()) };
1082 }
1083 let len = (Self::length(&a) + Self::length(&b)).saturating_sub(1);
1084 let block_len = if min_len >= 1 << 20 {
1085 1 << 20
1086 } else {
1087 (min_len.next_power_of_two() * 8).min(1 << 21) - min_len + 1
1088 };
1089 if max_len <= block_len {
1090 return convolve_u64_fft(a, b);
1091 }
1092 let mut result = vec![0u64; len];
1093 for (i, a) in a.chunks(block_len).enumerate() {
1094 for (j, b) in b.chunks(block_len).enumerate() {
1095 if a.len().min(b.len()) <= 60 {
1096 for (x, &a) in a.iter().enumerate() {
1097 for (y, &b) in b.iter().enumerate() {
1098 let value = &mut result[(i + j) * block_len + x + y];
1099 *value = value.wrapping_add(a.wrapping_mul(b));
1100 }
1101 }
1102 continue;
1103 }
1104 let product = convolve_u64_fft(a.to_vec(), b.to_vec());
1105 for (value, product) in result[(i + j) * block_len..].iter_mut().zip(product) {
1106 *value = value.wrapping_add(product);
1107 }
1108 }
1109 }
1110 result
1111 }
1112}
1113
1114fn convolve_u64_fft(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1115 crate::avx_helper!(@dispatch_avx2_fma return unsafe {
1117 convolve_u64_avx2(a, b)
1118 }, ());
1119 convolve_u64_fft_scalar(a, b)
1120}
1121
1122fn convolve_u64_fft_scalar(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
1123 fn split(values: &[u64]) -> [Vec<i64>; 5] {
1124 let mut result = std::array::from_fn(|_| Vec::with_capacity(values.len()));
1125 for mut value in values.iter().copied() {
1126 for part in &mut result {
1127 let digit = ((value << 51) as i64) >> 51;
1128 part.push(digit);
1129 value = (value >> 13).wrapping_add(u64::from(digit < 0));
1130 }
1131 }
1132 result
1133 }
1134
1135 let len = a.len() + b.len() - 1;
1136 let transform = |values: &[u64]| {
1137 if values.iter().any(|&value| value > u32::MAX as u64) {
1138 return split(values).map(|part| ConvolveRealFft::transform(part, len));
1139 }
1140 let [a, b, c, _, _] = split(values);
1141 let a = ConvolveRealFft::transform(a, len);
1142 let size = a.len();
1143 [
1144 a,
1145 ConvolveRealFft::transform(b, len),
1146 ConvolveRealFft::transform(c, len),
1147 vec![Zero::zero(); size],
1148 vec![Zero::zero(); size],
1149 ]
1150 };
1151 let fa = transform(&a);
1152 drop(a);
1153 let fb = transform(&b);
1154 drop(b);
1155 let values: [Vec<i64>; 5] = std::array::from_fn(|part| {
1156 let mut sum = fa[0].clone();
1157 ConvolveRealFft::multiply(&mut sum, &fb[part]);
1158 for left in 1..=part {
1159 let mut product = fa[left].clone();
1160 ConvolveRealFft::multiply(&mut product, &fb[part - left]);
1161 for (sum, product) in sum.iter_mut().zip(product) {
1162 *sum += product;
1163 }
1164 }
1165 ConvolveRealFft::inverse_transform(sum, len)
1166 });
1167 (0..len)
1168 .map(|i| {
1169 (values[0][i] as u64)
1170 .wrapping_add((values[1][i] as u64) << 13)
1171 .wrapping_add((values[2][i] as u64) << 26)
1172 .wrapping_add((values[3][i] as u64) << 39)
1173 .wrapping_add((values[4][i] as u64) << 52)
1174 })
1175 .collect()
1176}
1177
1178pub trait NttReuse: ConvolveSteps {
1179 const MULTIPLE: bool = true;
1180
1181 fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1183 Self::transform(t, len)
1184 }
1185
1186 fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1188 Self::inverse_transform(f, len)
1189 }
1190
1191 fn ntt_doubling(f: Self::F, monic: bool) -> Self::F;
1195
1196 fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1198
1199 fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1201
1202 fn multiply_prefix(f: &mut Self::F, g: &Self::F);
1204
1205 fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F);
1207
1208 fn max_product_sum_count(_f: &Self::F) -> usize {
1212 if Self::MULTIPLE { 1 } else { usize::MAX }
1213 }
1214
1215 fn power_projection_step(
1216 p_flat: Self::T,
1217 q_flat: Self::T,
1218 n: usize,
1219 py: usize,
1220 qy: usize,
1221 ) -> (Self::T, Self::T) {
1222 let base = n * 2;
1223 let len_p = base * py;
1224 let len_q = base * qy;
1225 let len = (len_p + len_q - 1).max(len_q + len_q - 1);
1226 let half = len.max(1).next_power_of_two() / 2;
1227
1228 let p_fft = Self::transform_ntt(p_flat, len);
1229 let q_fft = Self::transform_ntt(q_flat, len);
1230 let pr_fft = Self::odd_mul_normal_neg(&p_fft, &q_fft);
1231 let qr_fft = Self::even_mul_normal_neg(&q_fft, &q_fft);
1232 (
1233 Self::inverse_transform_ntt(pr_fft, half),
1234 Self::inverse_transform_ntt(qr_fft, half),
1235 )
1236 }
1237}
1238
1239thread_local!(
1240 static BIT_REVERSE: UnsafeCell<Vec<Vec<usize>>> = const { UnsafeCell::new(vec![]) };
1241);
1242
1243impl<M> NttReuse for Convolve<M>
1244where
1245 M: Montgomery32NttModulus,
1246{
1247 const MULTIPLE: bool = false;
1248
1249 fn transform_ntt(mut t: Self::T, len: usize) -> Self::F {
1250 t.resize_with(len.max(1).next_power_of_two(), Zero::zero);
1251 ntt(&mut t);
1252 t
1253 }
1254
1255 fn inverse_transform_ntt(mut f: Self::F, len: usize) -> Self::T {
1256 intt(&mut f);
1257 f.truncate(len);
1258 f
1259 }
1260
1261 fn ntt_doubling(mut f: Self::F, monic: bool) -> Self::F {
1262 let n = f.len();
1263 let k = n.trailing_zeros() as usize;
1264 let mut a = Self::inverse_transform_ntt(f.clone(), n);
1265 if monic {
1266 a[0] -= MInt::<M>::from(2);
1267 }
1268 let zeta = MInt::<M>::new_unchecked(M::INFO.root[k + 1]);
1269 let zeta2 = zeta * zeta;
1270 let mut rot = [MInt::one(), zeta, zeta2, zeta2 * zeta];
1271 let step = zeta2 * zeta2;
1272 for a in a.chunks_mut(4) {
1273 for (a, rot) in a.iter_mut().zip(&mut rot) {
1274 *a *= *rot;
1275 *rot *= step;
1276 }
1277 }
1278 f.extend(Self::transform_ntt(a, n));
1279 f
1280 }
1281
1282 fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1283 assert_eq!(f.len(), g.len());
1284 assert!(f.len().is_power_of_two());
1285 assert!(f.len() >= 2);
1286 if std::ptr::eq(f, g) {
1287 return f.as_chunks::<2>().0.iter().map(|a| a[0] * a[1]).collect();
1288 }
1289 let inv2 = MInt::<M>::from(2).inv();
1290 let n = f.len() / 2;
1291 (0..n)
1292 .map(|i| (f[i << 1] * g[i << 1 | 1] + f[i << 1 | 1] * g[i << 1]) * inv2)
1293 .collect()
1294 }
1295
1296 fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1297 assert_eq!(f.len(), g.len());
1298 assert!(f.len().is_power_of_two());
1299 assert!(f.len() >= 2);
1300 let mut inv2 = MInt::<M>::from(2).inv();
1301 let n = f.len() / 2;
1302 let k = f.len().trailing_zeros() as usize;
1303 let mut h = vec![MInt::<M>::zero(); n];
1304 let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1305 BIT_REVERSE.with(|br| {
1306 let br = unsafe { &mut *br.get() };
1307 if br.len() < k {
1308 br.resize_with(k, Default::default);
1309 }
1310 let k = k - 1;
1311 if br[k].is_empty() {
1312 let mut v = vec![0; 1 << k];
1313 for i in 0..1 << k {
1314 v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1315 }
1316 br[k] = v;
1317 }
1318 for &i in &br[k] {
1319 h[i] = (f[i << 1] * g[i << 1 | 1] - f[i << 1 | 1] * g[i << 1]) * inv2;
1320 inv2 *= w;
1321 }
1322 });
1323 h
1324 }
1325
1326 fn multiply_prefix(f: &mut Self::F, g: &Self::F) {
1327 pointwise_multiply(f, g);
1328 }
1329
1330 fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F) {
1331 assert!(sum.len() == f.len() && sum.len() == g.len());
1332 pointwise_multiply_add(sum, f, g);
1333 }
1334
1335 fn power_projection_step(
1336 p_flat: Vec<MInt<M>>,
1337 q_flat: Vec<MInt<M>>,
1338 n: usize,
1339 py: usize,
1340 qy: usize,
1341 ) -> (Vec<MInt<M>>, Vec<MInt<M>>) {
1342 let high_degree = (qy - 1) * 2;
1343 let rows = (py + qy - 1).max(high_degree).next_power_of_two();
1344 let cols = n * 2;
1345 let size = rows * cols;
1346 let mut p = p_flat;
1347 p.resize_with(size, MInt::<M>::zero);
1348 ntt_rows(&mut p, cols);
1349 ntt_batch(&mut p, cols);
1350
1351 let mut q = q_flat;
1352 q.resize_with(size, MInt::<M>::zero);
1353 ntt_rows(&mut q, cols);
1354 let q_high = (rows == high_degree).then(|| q[(qy - 1) * cols..qy * cols].to_vec());
1355 ntt_batch(&mut q, cols);
1356
1357 let half = cols / 2;
1358 let mut odd_factor = vec![MInt::<M>::zero(); half];
1359 let mut factor = MInt::<M>::from(2).inv();
1360 let k = cols.trailing_zeros() as usize;
1361 let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1362 BIT_REVERSE.with(|br| {
1363 let br = unsafe { &mut *br.get() };
1364 if br.len() < k {
1365 br.resize_with(k, Default::default);
1366 }
1367 let k = k - 1;
1368 if br[k].is_empty() {
1369 let mut v = vec![0; 1 << k];
1370 for i in 0..1 << k {
1371 v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1372 }
1373 br[k] = v;
1374 }
1375 for &i in &br[k] {
1376 odd_factor[i] = factor;
1377 factor *= w;
1378 }
1379 });
1380
1381 let mut pr = vec![MInt::<M>::zero(); rows * half];
1382 let mut qr = vec![MInt::<M>::zero(); rows * half];
1383 for i in 0..pr.len() {
1384 pr[i] = (p[i << 1] * q[i << 1 | 1] - p[i << 1 | 1] * q[i << 1])
1385 * odd_factor[i & (half - 1)];
1386 qr[i] = q[i << 1] * q[i << 1 | 1];
1387 }
1388 intt_batch(&mut pr, half);
1389 intt_rows(&mut pr, half);
1390 intt_batch(&mut qr, half);
1391 intt_rows(&mut qr, half);
1392
1393 if let Some(q_high) = q_high {
1394 let mut q_high_even = vec![MInt::<M>::zero(); half];
1395 for i in 0..half {
1396 q_high_even[i] = q_high[i << 1] * q_high[i << 1 | 1];
1397 }
1398 intt(&mut q_high_even);
1399 for (value, high) in qr.iter_mut().zip(&q_high_even) {
1400 *value -= *high;
1401 }
1402 qr.extend_from_slice(&q_high_even);
1403 }
1404 (pr, qr)
1405 }
1406}
1407
1408impl<M, N1, N2, N3> NttReuse for Convolve<(M, (N1, N2, N3))>
1409where
1410 M: MIntConvert + MIntConvert<u32>,
1411 N1: Montgomery32NttModulus,
1412 N2: Montgomery32NttModulus,
1413 N3: Montgomery32NttModulus,
1414{
1415 fn max_product_sum_count(f: &Self::F) -> usize {
1416 let modulus = <M as MIntConvert<u32>>::mod_into() as u128;
1417 if modulus == 1 {
1418 return usize::MAX;
1419 }
1420 let capacity = N1::MOD as u128 * N2::MOD as u128 * N3::MOD as u128;
1421 ((capacity - 1) / ((modulus - 1) * (modulus - 1)) / f.0.len() as u128)
1422 .clamp(1, usize::MAX as u128) as usize
1423 }
1424
1425 fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1426 let npot = len.max(1).next_power_of_two();
1427 let f = convert_crt_input(t, npot);
1428 (
1429 Convolve::<N1>::transform_ntt(f.0, npot),
1430 Convolve::<N2>::transform_ntt(f.1, npot),
1431 Convolve::<N3>::transform_ntt(f.2, npot),
1432 )
1433 }
1434
1435 fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1436 reconstruct_mint_crt((
1437 Convolve::<N1>::inverse_transform_ntt(f.0, len),
1438 Convolve::<N2>::inverse_transform_ntt(f.1, len),
1439 Convolve::<N3>::inverse_transform_ntt(f.2, len),
1440 ))
1441 }
1442
1443 fn ntt_doubling(f: Self::F, monic: bool) -> Self::F {
1444 if monic {
1445 let n = f.0.len();
1446 let mut coefficients = Self::inverse_transform_ntt(f, n);
1447 coefficients[0] -= MInt::<M>::one();
1448 coefficients.push(MInt::<M>::one());
1449 Self::transform_ntt(coefficients, n * 2)
1450 } else {
1451 (
1452 Convolve::<N1>::ntt_doubling(f.0, false),
1453 Convolve::<N2>::ntt_doubling(f.1, false),
1454 Convolve::<N3>::ntt_doubling(f.2, false),
1455 )
1456 }
1457 }
1458
1459 fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1460 fn even_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1461 where
1462 M: Montgomery32NttModulus,
1463 {
1464 let n = f.len();
1465 assert_eq!(f.len(), g.len());
1466 assert!(f.len().is_power_of_two());
1467 assert!(f.len() >= 2);
1468 let inv2 = MInt::<M>::from(2).inv();
1469 let u = MInt::<M>::new(m) * MInt::<M>::from(n as u32);
1470 let n = f.len() / 2;
1471 (0..n)
1472 .map(|i| {
1473 (f[i << 1]
1474 * if i == 0 {
1475 g[i << 1 | 1] + u
1476 } else {
1477 g[i << 1 | 1]
1478 }
1479 + f[i << 1 | 1] * g[i << 1])
1480 * inv2
1481 })
1482 .collect()
1483 }
1484
1485 let m = M::mod_into();
1486 (
1487 even_mul_normal_neg_corrected(&f.0, &g.0, m),
1488 even_mul_normal_neg_corrected(&f.1, &g.1, m),
1489 even_mul_normal_neg_corrected(&f.2, &g.2, m),
1490 )
1491 }
1492
1493 fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F {
1494 fn odd_mul_normal_neg_corrected<M>(f: &[MInt<M>], g: &[MInt<M>], m: u32) -> Vec<MInt<M>>
1495 where
1496 M: Montgomery32NttModulus,
1497 {
1498 assert_eq!(f.len(), g.len());
1499 assert!(f.len().is_power_of_two());
1500 assert!(f.len() >= 2);
1501 let mut inv2 = MInt::<M>::from(2).inv();
1502 let u = MInt::<M>::new(m) * MInt::<M>::from(f.len() as u32);
1503 let n = f.len() / 2;
1504 let k = f.len().trailing_zeros() as usize;
1505 let mut h = vec![MInt::<M>::zero(); n];
1506 let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1507 BIT_REVERSE.with(|br| {
1508 let br = unsafe { &mut *br.get() };
1509 if br.len() < k {
1510 br.resize_with(k, Default::default);
1511 }
1512 let k = k - 1;
1513 if br[k].is_empty() {
1514 let mut v = vec![0; 1 << k];
1515 for i in 0..1 << k {
1516 v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1517 }
1518 br[k] = v;
1519 }
1520 for &i in &br[k] {
1521 h[i] = (f[i << 1]
1522 * if i == 0 {
1523 g[i << 1 | 1] + u
1524 } else {
1525 g[i << 1 | 1]
1526 }
1527 - f[i << 1 | 1] * g[i << 1])
1528 * inv2;
1529 inv2 *= w;
1530 }
1531 });
1532 h
1533 }
1534
1535 let m = M::mod_into();
1536 (
1537 odd_mul_normal_neg_corrected(&f.0, &g.0, m),
1538 odd_mul_normal_neg_corrected(&f.1, &g.1, m),
1539 odd_mul_normal_neg_corrected(&f.2, &g.2, m),
1540 )
1541 }
1542
1543 fn multiply_prefix(f: &mut Self::F, g: &Self::F) {
1544 Convolve::<N1>::multiply_prefix(&mut f.0, &g.0);
1545 Convolve::<N2>::multiply_prefix(&mut f.1, &g.1);
1546 Convolve::<N3>::multiply_prefix(&mut f.2, &g.2);
1547 }
1548
1549 fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F) {
1550 Convolve::<N1>::multiply_add(&mut sum.0, &f.0, &g.0);
1551 Convolve::<N2>::multiply_add(&mut sum.1, &f.1, &g.1);
1552 Convolve::<N3>::multiply_add(&mut sum.2, &f.2, &g.2);
1553 }
1554}
1555
1556#[cfg(test)]
1557mod tests {
1558 use super::*;
1559 use crate::num::{mint_basic::Modulo1000000009, montgomery::MInt998244353};
1560 use crate::tools::Xorshift;
1561 #[cfg(target_arch = "x86_64")]
1562 use crate::tools::avx512_supported;
1563
1564 #[test]
1565 fn test_ntt_batch() {
1566 fn check<M: Montgomery32NttModulus>() {
1567 let mut rng = Xorshift::default();
1568 for log_n in 0..=8 {
1569 let n = 1 << log_n;
1570 for width in 1..=64 {
1571 let input: Vec<MInt<M>> = rng.random_iter(..).take(n * width).collect();
1572 let mut expected = input.clone();
1573 ntt_batch_scalar(&mut expected, width);
1574
1575 let mut actual = input.clone();
1576 ntt_batch(&mut actual, width);
1577 assert_eq!(actual, expected);
1578 intt_batch(&mut actual, width);
1579 assert_eq!(actual, input);
1580
1581 #[cfg(target_arch = "x86_64")]
1582 if is_x86_feature_detected!("avx2") {
1583 let mut actual = input.clone();
1584 unsafe { ntt_simd::ntt_batch_avx2::<_, true>(&mut actual, width) };
1585 assert_eq!(actual, expected);
1586 unsafe { ntt_simd::intt_batch_avx2::<_, true>(&mut actual, width) };
1587 assert_eq!(actual, input);
1588 }
1589
1590 #[cfg(target_arch = "x86_64")]
1591 if avx512_supported() {
1592 let mut actual = input.clone();
1593 unsafe { ntt_simd::ntt_batch_avx512(&mut actual, width) };
1594 assert_eq!(actual, expected);
1595 unsafe { ntt_simd::intt_batch_avx512::<_, false>(&mut actual, width) };
1596 assert_eq!(actual, input);
1597 }
1598 }
1599 }
1600 }
1601
1602 enum Modulo2013265921 {}
1603 impl MontgomeryReduction32 for Modulo2013265921 {
1604 const MOD: u32 = 2013265921;
1605 }
1606 impl Montgomery32NttModulus for Modulo2013265921 {
1607 const PRIMITIVE_ROOT: u32 = 31;
1608 }
1609 check::<Modulo998244353>();
1610 check::<Modulo2013265921>();
1611 }
1612
1613 #[test]
1614 fn test_convolve_naive() {
1615 let mut rng = Xorshift::default();
1616 for case in 0..1030 {
1617 let (n, m) = if case < 1000 {
1618 (rng.random(0..=60), rng.random(0..=60))
1619 } else {
1620 (
1621 (case / 3 % 3 + 1) * 1024 + case % 3 - 1,
1622 rng.random(0..=if case < 1015 { 64 } else { 1025 }),
1623 )
1624 };
1625 let a: Vec<u32> = rng.random_iter(0u32..1000).take(n).collect();
1626 let b: Vec<u32> = rng.random_iter(0u32..1000).take(m).collect();
1627 let mut c = vec![0u32; (n + m).saturating_sub(1)];
1628 for i in 0..n {
1629 for j in 0..m {
1630 c[i + j] += a[i] * b[j];
1631 }
1632 }
1633 assert_eq!(c, convolve_naive(&a, &b));
1634 assert_eq!(c, convolve_naive(&b, &a));
1635 }
1636 }
1637
1638 #[test]
1639 fn test_convolve_karatsuba() {
1640 let mut rng = Xorshift::default();
1641 for _ in 0..1000 {
1642 let n = if rng.gen_bool(0.1) {
1643 rng.random(201..=4096)
1644 } else {
1645 rng.random(0..=200)
1646 };
1647 let m = rng.random(0..=200);
1648 let a: Vec<u32> = rng.random_iter(0u32..1000).take(n).collect();
1649 let b: Vec<u32> = rng.random_iter(0u32..1000).take(m).collect();
1650 let mut c = vec![0u32; (n + m).saturating_sub(1)];
1651 for i in 0..n {
1652 for j in 0..m {
1653 c[i + j] += a[i] * b[j];
1654 }
1655 }
1656 let d = convolve_karatsuba(&a, &b);
1657 assert_eq!(c, d);
1658 assert_eq!(c, convolve_karatsuba(&b, &a));
1659 }
1660 }
1661
1662 #[test]
1663 fn test_ntt998244353() {
1664 let mut rng = Xorshift::default();
1665 for _ in 0..1000 {
1666 let (n, m) = if rng.random(0..100) == 0 {
1667 let w = rng.random(6..=8);
1668 ((1usize << w) + 1usize, (1usize << w) + 1usize)
1669 } else {
1670 let n = rng.random(0..=5);
1671 let m = rng.random(0..=5);
1672 (
1673 if n == 5 { rng.random(70..=120) } else { n },
1674 if m == 5 { rng.random(70..=120) } else { m },
1675 )
1676 };
1677 let a: Vec<MInt998244353> = rng.random_iter(..).take(n).collect();
1678 let mut b: Vec<MInt998244353> = rng.random_iter(..).take(m).collect();
1679 if n == m && rng.random(0..2) == 0 {
1680 b = a.clone();
1681 }
1682
1683 let mut c = vec![MInt998244353::zero(); (n + m).saturating_sub(1)];
1684 for i in 0..n {
1685 for j in 0..m {
1686 c[i + j] += a[i] * b[j];
1687 }
1688 }
1689 let d = Convolve998244353::convolve(a, b);
1690 assert_eq!(c, d);
1691 }
1692 assert_eq!(NttInfo::new::<Modulo998244353>(), Modulo998244353::INFO);
1693 }
1694
1695 #[test]
1696 fn test_convolve_large_ntt() {
1697 enum Modulo17 {}
1698 impl MontgomeryReduction32 for Modulo17 {
1699 const MOD: u32 = 17;
1700 }
1701 impl Montgomery32NttModulus for Modulo17 {}
1702 enum Modulo97 {}
1703 impl MontgomeryReduction32 for Modulo97 {
1704 const MOD: u32 = 97;
1705 }
1706 impl Montgomery32NttModulus for Modulo97 {
1707 const PRIMITIVE_ROOT: u32 = 5;
1708 }
1709 enum Modulo193 {}
1710 impl MontgomeryReduction32 for Modulo193 {
1711 const MOD: u32 = 193;
1712 }
1713 impl Montgomery32NttModulus for Modulo193 {
1714 const PRIMITIVE_ROOT: u32 = 5;
1715 }
1716
1717 fn check<M: Montgomery32NttModulus>() {
1718 let mut rng = Xorshift::default();
1719 for multiplier in [1, 2, 4, 8, 16] {
1720 let cap = multiplier << M::RANK;
1721 for _ in 0..40 {
1722 let len = rng.random(cap / 2..=cap + 2);
1723 let n = rng.random(0..=len);
1724 let m = len - n;
1725 let a: Vec<MInt<M>> = rng
1726 .random_iter(0u32..M::MOD)
1727 .take(n)
1728 .map(MInt::from)
1729 .collect();
1730 let b = if rng.gen_bool(0.25) {
1731 a.clone()
1732 } else {
1733 rng.random_iter(0u32..M::MOD)
1734 .take(m)
1735 .map(MInt::from)
1736 .collect()
1737 };
1738 assert_eq!(convolve_naive(&a, &b), Convolve::<M>::convolve(a, b));
1739 }
1740 }
1741 }
1742 check::<Modulo17>();
1743 check::<Modulo97>();
1744 check::<Modulo193>();
1745 }
1746
1747 #[test]
1748 fn test_convolve3() {
1749 use crate::num::mint_basic::{DynModuloU32, Modulo2};
1750
1751 fn check<M: MIntConvert<u32> + MIntBase<Inner = u32>>() {
1752 let modulus = M::get_mod();
1753 let mut rng = Xorshift::default();
1754 for case in 0..1000 {
1755 let n = rng.random(0..=5);
1756 let n = if case == 0 {
1757 rng.random(8192..=12288)
1758 } else if n == 5 {
1759 rng.random(5..=600)
1760 } else {
1761 n
1762 };
1763 let m = rng.random(0..=5);
1764 let m = if case == 0 {
1765 rng.random(257..=512)
1766 } else if m == 5 {
1767 rng.random(5..=600)
1768 } else {
1769 m
1770 };
1771 let a: Vec<u32> = rng.random_iter(0..modulus).take(n).collect();
1772 let b: Vec<u32> = rng.random_iter(0..modulus).take(m).collect();
1773 let mut expected = vec![0u128; (n + m).saturating_sub(1)];
1774 for (i, &a) in a.iter().enumerate() {
1775 for (j, &b) in b.iter().enumerate() {
1776 expected[i + j] += a as u128 * b as u128;
1777 }
1778 }
1779 let expected: Vec<_> = expected
1780 .into_iter()
1781 .map(|x| (x % modulus as u128) as u32)
1782 .collect();
1783 let actual = MIntConvolve::<M>::convolve(
1784 a.into_iter().map(MInt::from).collect(),
1785 b.into_iter().map(MInt::from).collect(),
1786 );
1787 assert_eq!(
1788 actual.into_iter().map(u32::from).collect::<Vec<_>>(),
1789 expected
1790 );
1791 }
1792 }
1793 check::<Modulo1000000009>();
1794 check::<Modulo2>();
1795 let mut rng = Xorshift::default();
1796 for modulus in [1, u32::MAX]
1797 .into_iter()
1798 .chain(rng.random_iter(1..).take(8))
1799 {
1800 DynModuloU32::set_mod(modulus);
1801 check::<DynModuloU32>();
1802 }
1803 enum Modulo<const M: u32> {}
1804 impl<const M: u32> MontgomeryReduction32 for Modulo<M> {
1805 const MOD: u32 = M;
1806 }
1807 impl<const M: u32> Montgomery32NttModulus for Modulo<M> {}
1808 for modulus in [1499, u32::MAX] {
1809 DynModuloU32::set_mod(modulus);
1810 for _ in 0..40 {
1811 let n = rng.random(250..400);
1812 let m = rng.random(250..400);
1813 let a: Vec<u32> = rng.random_iter(modulus - 9..modulus).take(n).collect();
1814 let b: Vec<u32> = rng.random_iter(modulus - 9..modulus).take(m).collect();
1815 let mut expected = vec![0u128; n + m - 1];
1816 for (i, &a) in a.iter().enumerate() {
1817 for (j, &b) in b.iter().enumerate() {
1818 expected[i + j] += a as u128 * b as u128;
1819 }
1820 }
1821 let actual =
1822 convolve_mint_crt::<DynModuloU32, Modulo<257>, Modulo<769>, Modulo<3329>>(
1823 a.into_iter().map(MInt::from).collect(),
1824 b.into_iter().map(MInt::from).collect(),
1825 );
1826 assert_eq!(actual.len(), expected.len());
1827 for (actual, expected) in actual.into_iter().zip(expected) {
1828 assert_eq!(u32::from(actual), (expected % modulus as u128) as u32);
1829 }
1830 }
1831 }
1832 DynModuloU32::set_mod(1_000_000_007);
1833 }
1834
1835 #[test]
1836 fn test_convolve3_large_coefficients() {
1837 use crate::num::mint_basic::{DynMIntU32, DynModuloU32};
1838 let mut rng = Xorshift::default();
1839 for _ in 0..3 {
1840 let modulus = u32::MAX - rng.random(0u32..65536);
1841 DynModuloU32::set_mod(modulus);
1842 let n = (1 << 18) - rng.random(0usize..1024);
1843 let m = (1 << 18) - rng.random(0usize..1024);
1844 let x = modulus / 2 - rng.random(32700u32..32800);
1845 let y = modulus / 2 - rng.random(32700u32..32800);
1846 let actual = MIntConvolve::<DynModuloU32>::convolve(
1847 vec![DynMIntU32::from(x); n],
1848 vec![DynMIntU32::from(y); m],
1849 );
1850 assert_eq!(actual.len(), n + m - 1);
1851 for (i, actual) in actual.into_iter().enumerate() {
1852 let count = (i + 1).min(n).min(m).min(n + m - 1 - i);
1853 let expected = (count as u128 * x as u128 * y as u128 % modulus as u128) as u32;
1854 assert_eq!(u32::from(actual), expected);
1855 }
1856 }
1857 for (modulus, log_n) in [(1_000_000_007, 19), (u32::MAX, 17)] {
1858 DynModuloU32::set_mod(modulus);
1859 let n = (1 << log_n) - rng.random(0usize..1024);
1860 let m = rng.random(n / 2..=n);
1861 let a: Vec<u32> = rng.random_iter(0..modulus).take(n).collect();
1862 let y = rng.random(1..modulus);
1863 let actual = MIntConvolve::<DynModuloU32>::convolve(
1864 a.iter().copied().map(DynMIntU32::from).collect(),
1865 vec![DynMIntU32::from(y); m],
1866 );
1867 assert_eq!(actual.len(), n + m - 1);
1868 let mut sum = 0u128;
1869 for (i, actual) in actual.into_iter().enumerate() {
1870 if i < n {
1871 sum += a[i] as u128;
1872 }
1873 if i >= m && i - m < n {
1874 sum -= a[i - m] as u128;
1875 }
1876 assert_eq!(
1877 u32::from(actual),
1878 (sum * y as u128 % modulus as u128) as u32
1879 );
1880 }
1881 }
1882 DynModuloU32::set_mod(1_000_000_007);
1883 }
1884
1885 #[test]
1886 fn test_convolve_u64() {
1887 enum Modulo97 {}
1888 impl MontgomeryReduction32 for Modulo97 {
1889 const MOD: u32 = 97;
1890 }
1891 impl Montgomery32NttModulus for Modulo97 {}
1892 type SmallCrt = Convolve<(u64, (Modulo998244353, Modulo469762049, Modulo97))>;
1893 let mut rng = Xorshift::default();
1894 for case in 0..1000 {
1895 let (n, m) = if case < 36 {
1896 (case / 6, case % 6)
1897 } else if rng.gen_bool(0.01) {
1898 (rng.random(1537..=2000), rng.random(513..=800))
1899 } else {
1900 (rng.random(0..=400), rng.random(0..=400))
1901 };
1902 let mask = if rng.gen_bool(0.5) {
1903 u32::MAX as u64
1904 } else {
1905 u64::MAX
1906 };
1907 let a: Vec<u64> = rng.random_iter(..).map(|a: u64| a & mask).take(n).collect();
1908 let mask = if rng.gen_bool(0.5) {
1909 u32::MAX as u64
1910 } else {
1911 u64::MAX
1912 };
1913 let b: Vec<u64> = rng.random_iter(..).map(|b: u64| b & mask).take(m).collect();
1914 let mut c = vec![0u64; (n + m).saturating_sub(1)];
1915 for i in 0..n {
1916 for j in 0..m {
1917 c[i + j] = c[i + j].wrapping_add(a[i].wrapping_mul(b[j]));
1918 }
1919 }
1920 let mut f = U64Convolve::transform(a.clone(), c.len());
1921 let g = U64Convolve::transform(b.clone(), c.len());
1922 U64Convolve::multiply(&mut f, &g);
1923 assert_eq!(U64Convolve::inverse_transform(f, c.len()), c);
1924 if c.len() <= 32 {
1925 let mut f = SmallCrt::transform(a.clone(), c.len());
1926 let g = SmallCrt::transform(b.clone(), c.len());
1927 SmallCrt::multiply(&mut f, &g);
1928 assert_eq!(SmallCrt::inverse_transform(f, c.len()), c);
1929 }
1930 assert_eq!(U64Convolve::convolve(a.clone(), b), c);
1931 let f = U64Convolve::transform(a.clone(), n);
1932 assert_eq!(U64Convolve::inverse_transform(f, n), a);
1933 let mut square = vec![0u64; (n * 2).saturating_sub(1)];
1934 for (i, &x) in a.iter().enumerate() {
1935 for (j, &y) in a.iter().enumerate() {
1936 square[i + j] = square[i + j].wrapping_add(x.wrapping_mul(y));
1937 }
1938 }
1939 assert_eq!(U64Convolve::square(a, square.len()), square);
1940 }
1941
1942 for shift in [12, 15] {
1943 let n = (1 << 19) + rng.random(1usize..1024);
1944 let m = (1 << 19) + rng.random(1usize..1024);
1945 let x = (rng.rand64() | 1) << shift;
1946 let y = (rng.rand64() | 1) << shift;
1947 let alternating = rng.gen_bool(0.5);
1948 let a = (0..n)
1949 .map(|i| {
1950 if alternating && i & 1 == 1 {
1951 x.wrapping_neg()
1952 } else {
1953 x
1954 }
1955 })
1956 .collect();
1957 let b = (0..m)
1958 .map(|i| {
1959 if alternating && i & 1 == 1 {
1960 y.wrapping_neg()
1961 } else {
1962 y
1963 }
1964 })
1965 .collect();
1966 let actual = U64Convolve::convolve(a, b);
1967 assert_eq!(actual.len(), n + m - 1);
1968 for (i, actual) in actual.into_iter().enumerate() {
1969 let count = (i + 1).min(n).min(m).min(n + m - 1 - i) as u64;
1970 let expected = count.wrapping_mul(x).wrapping_mul(y);
1971 let expected = if alternating && i & 1 == 1 {
1972 expected.wrapping_neg()
1973 } else {
1974 expected
1975 };
1976 assert_eq!(actual, expected, "{n}x{m}/{shift}/{i}");
1977 }
1978 }
1979 let n = (1 << 21) + rng.random(1usize..1024);
1980 let m = rng.random(1537usize..4096);
1981 let x = rng.rand64();
1982 let y = rng.rand64();
1983 let actual = U64Convolve::convolve(vec![x; n], vec![y; m]);
1984 assert_eq!(actual.len(), n + m - 1);
1985 for (i, actual) in actual.into_iter().enumerate() {
1986 let count = (i + 1).min(n).min(m).min(n + m - 1 - i) as u64;
1987 assert_eq!(actual, count.wrapping_mul(x).wrapping_mul(y));
1988 }
1989 }
1990
1991 #[test]
1992 fn test_ntt_reuse_998244353() {
1993 let mut rng = Xorshift::default();
1994 for _ in 0..100 {
1995 let n: usize = if rng.gen_bool(0.5) {
1996 rng.random(1..=20)
1997 } else {
1998 rng.random(1..=1000)
1999 };
2000 let a: Vec<MInt998244353> = rng.random_iter(..).take(n).collect();
2001 let f = Convolve998244353::transform_ntt(a.clone(), n);
2002
2003 {
2005 for constant in [
2006 MInt998244353::zero(),
2007 MInt998244353::one(),
2008 -MInt998244353::one(),
2009 a[0],
2010 ] {
2011 let mut cyclic = a.clone();
2012 cyclic[0] = constant + MInt998244353::one();
2013 let f_monic = Convolve998244353::transform_ntt(cyclic, n);
2014 let f_monic = Convolve998244353::ntt_doubling(f_monic, true);
2015 let mut monic = a.clone();
2016 monic[0] = constant;
2017 monic.resize_with(n.next_power_of_two(), Zero::zero);
2018 monic.push(MInt998244353::one());
2019 assert_eq!(f_monic, Convolve998244353::transform_ntt(monic, n * 2));
2020 }
2021 let f_double = Convolve998244353::ntt_doubling(f.clone(), false);
2022 let mut a = a.clone();
2023 a.resize_with(n * 2, Zero::zero);
2024 assert_eq!(f_double, Convolve998244353::transform_ntt(a, n * 2));
2025 }
2026
2027 let f = Convolve998244353::transform_ntt(a.clone(), n * 2);
2028 let b: Vec<MInt998244353> = rng.random_iter(..).take(n).collect();
2029 let g = Convolve998244353::transform_ntt(b.clone(), n * 2);
2030 let mut b_neg = b.clone();
2031 for b in b_neg.iter_mut().skip(1).step_by(2) {
2032 *b = -*b;
2033 }
2034
2035 {
2037 let fg_neg = Convolve998244353::even_mul_normal_neg(&f, &g);
2038 let ab_neg_even: Vec<_> = Convolve998244353::convolve(a.clone(), b_neg.clone())
2039 .into_iter()
2040 .step_by(2)
2041 .collect();
2042 assert_eq!(fg_neg, Convolve998244353::transform_ntt(ab_neg_even, n));
2043 }
2044
2045 {
2047 let fg_neg = Convolve998244353::odd_mul_normal_neg(&f, &g);
2048 let ab_neg_odd: Vec<_> = Convolve998244353::convolve(a.clone(), b_neg.clone())
2049 .into_iter()
2050 .skip(1)
2051 .step_by(2)
2052 .collect();
2053 assert_eq!(fg_neg, Convolve998244353::transform_ntt(ab_neg_odd, n));
2054 }
2055 }
2056 }
2057
2058 #[test]
2059 fn test_ntt_reuse_triple() {
2060 type M = MInt<Modulo1000000009>;
2061 let mut rng = Xorshift::default();
2062 for _ in 0..100 {
2063 let n: usize = if rng.gen_bool(0.5) {
2064 rng.random(1..=20)
2065 } else {
2066 rng.random(1..=1000)
2067 };
2068 let a: Vec<M> = rng.random_iter(..).take(n).collect();
2069 let f = MIntConvolve::<Modulo1000000009>::transform_ntt(a.clone(), n);
2070
2071 {
2073 for constant in [M::zero(), M::one(), -M::one(), a[0]] {
2074 let mut cyclic = a.clone();
2075 cyclic[0] = constant + M::one();
2076 let f_monic = MIntConvolve::<Modulo1000000009>::transform_ntt(cyclic, n);
2077 let f_monic = MIntConvolve::<Modulo1000000009>::ntt_doubling(f_monic, true);
2078 let mut monic = a.clone();
2079 monic[0] = constant;
2080 monic.resize_with(n.next_power_of_two(), Zero::zero);
2081 monic.push(M::one());
2082 assert_eq!(
2083 f_monic,
2084 MIntConvolve::<Modulo1000000009>::transform_ntt(monic, n * 2)
2085 );
2086 }
2087 let f_double = MIntConvolve::<Modulo1000000009>::ntt_doubling(f.clone(), false);
2088 let mut a = a.clone();
2089 a.resize_with(n * 2, Zero::zero);
2090 assert_eq!(
2091 f_double,
2092 MIntConvolve::<Modulo1000000009>::transform_ntt(a, n * 2)
2093 );
2094 }
2095
2096 let f = MIntConvolve::<Modulo1000000009>::transform_ntt(a.clone(), n * 2);
2097 let b: Vec<M> = rng.random_iter(..).take(n).collect();
2098 let g = MIntConvolve::<Modulo1000000009>::transform_ntt(b.clone(), n * 2);
2099 let mut b_neg = b.clone();
2100 for b in b_neg.iter_mut().skip(1).step_by(2) {
2101 *b = -*b;
2102 }
2103
2104 {
2106 let fg_neg = MIntConvolve::<Modulo1000000009>::even_mul_normal_neg(&f, &g);
2107 let ab_neg_even: Vec<_> =
2108 MIntConvolve::<Modulo1000000009>::convolve(a.clone(), b_neg.clone())
2109 .into_iter()
2110 .step_by(2)
2111 .collect();
2112 assert_eq!(
2113 MIntConvolve::<Modulo1000000009>::inverse_transform_ntt(fg_neg, n),
2114 ab_neg_even
2115 );
2116 }
2117
2118 {
2120 let fg_neg = MIntConvolve::<Modulo1000000009>::odd_mul_normal_neg(&f, &g);
2121 let ab_neg_odd: Vec<_> =
2122 MIntConvolve::<Modulo1000000009>::convolve(a.clone(), b_neg.clone())
2123 .into_iter()
2124 .skip(1)
2125 .step_by(2)
2126 .chain([M::zero()])
2127 .collect();
2128 assert_eq!(
2129 MIntConvolve::<Modulo1000000009>::inverse_transform_ntt(fg_neg, n),
2130 ab_neg_odd
2131 );
2132 }
2133 }
2134 }
2135 #[test]
2136 fn test_fps_crt_product_sum_capacity() {
2137 use crate::{
2138 math::FormalPowerSeries,
2139 num::mint_basic::{DynMIntU32 as Mint, DynModuloU32},
2140 };
2141 enum Modulo<const P: u32> {}
2142 impl<const P: u32> MontgomeryReduction32 for Modulo<P> {
2143 const MOD: u32 = P;
2144 }
2145 impl<const P: u32> Montgomery32NttModulus for Modulo<P> {}
2146 type C = Convolve<(DynModuloU32, (Modulo<257>, Modulo<769>, Modulo<3329>))>;
2147
2148 let mut rng = Xorshift::default();
2149 for modulus in [521, 2503] {
2150 Mint::set_mod(modulus);
2151 let degrees: Vec<_> = (6..=8)
2152 .flat_map(|k| (1 << k) - 1..=(1 << k) + 1)
2153 .chain((0..12).map(|_| rng.random(65..=300)))
2154 .collect();
2155 for deg in degrees {
2156 for random in [false, true] {
2157 let mut f = vec![Mint::zero(); deg];
2158 let mut power = Mint::one();
2159 for (i, value) in f.iter_mut().enumerate().skip(1) {
2160 power *= Mint::from(2);
2161 *value = if random {
2162 Mint::from(rng.random(0..modulus))
2163 } else {
2164 (Mint::one() - power) / Mint::from(i)
2165 };
2166 }
2167 let mut expected = vec![Mint::zero(); deg];
2168 expected[0] = Mint::one();
2169 for i in 1..deg {
2170 for j in 1..=i {
2171 let value = f[j] * Mint::from(j) * expected[i - j];
2172 expected[i] += value;
2173 }
2174 expected[i] /= Mint::from(i);
2175 }
2176 assert_eq!(
2177 FormalPowerSeries::<_, C>::from_vec(f).exp(deg).data,
2178 expected
2179 );
2180 let f = FormalPowerSeries::<_, C>::from_vec(expected);
2181 let mut expected = vec![Mint::zero(); deg];
2182 expected[0] = Mint::one();
2183 let rhs = rng.random(5..=8);
2184 for _ in 0..rhs {
2185 let mut next = vec![Mint::zero(); deg];
2186 for i in 0..deg {
2187 for j in 0..deg - i {
2188 next[i + j] += expected[i] * f[j];
2189 }
2190 }
2191 expected = next;
2192 }
2193 assert_eq!(f.pow(rhs, deg).data, expected);
2194 }
2195 }
2196 }
2197 Mint::set_mod(1);
2198 for log_n in 0..=8 {
2199 let n = 1 << log_n;
2200 let f = C::transform_ntt(vec![Mint::zero(); n], n);
2201 assert_eq!(C::max_product_sum_count(&f), usize::MAX);
2202 assert_eq!(C::inverse_transform_ntt(f, n), vec![Mint::zero(); n]);
2203 }
2204 }
2205
2206 #[test]
2207 fn test_crt_montgomery_coefficients() {
2208 let mut rng = Xorshift::default();
2209 let sizes: Vec<_> = (1..=8)
2210 .flat_map(|n| (1..=8).map(move |m| (n, m)))
2211 .chain((0..40).map(|_| (rng.random(280..=600), rng.random(280..=600))))
2212 .collect();
2213 for (n, m) in sizes {
2214 let a: Vec<MInt998244353> = rng.random_iter(..).take(n).collect();
2215 let b: Vec<MInt998244353> = rng.random_iter(..).take(m).collect();
2216 let f = MIntConvolve::<Modulo998244353>::transform(a.clone(), n);
2217 assert_eq!(MIntConvolve::<Modulo998244353>::inverse_transform(f, n), a);
2218 let f = MIntConvolve::<Modulo998244353>::transform_ntt(a.clone(), n);
2219 assert_eq!(
2220 MIntConvolve::<Modulo998244353>::inverse_transform_ntt(f, n),
2221 a
2222 );
2223 let mut expected = vec![0u64; n + m - 1];
2224 for (i, x) in a.iter().enumerate() {
2225 for (j, y) in b.iter().enumerate() {
2226 expected[i + j] =
2227 (expected[i + j] + x.inner() as u64 * y.inner() as u64) % 998244353;
2228 }
2229 }
2230 let actual = MIntConvolve::<Modulo998244353>::convolve(a, b);
2231 assert_eq!(
2232 actual
2233 .into_iter()
2234 .map(|x| x.inner() as u64)
2235 .collect::<Vec<_>>(),
2236 expected
2237 );
2238 }
2239 }
2240}