Skip to main content

competitive/math/
number_theoretic_transform.rs

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    // power_projection uses power-of-two widths; SIMD scalar tails regress for some other widths.
22    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>;
37/// Raw transforms require each integer coefficient reconstructed by CRT to be below
38/// the product of the three NTT moduli. `convolve` splits products exceeding this bound.
39pub type MIntConvolve<M> = Convolve<(M, (Modulo167772161, Modulo469762049, Modulo754974721))>;
40/// Convolution modulo 2^64. Multiply only freshly transformed operands; reconstruct
41/// and transform again before multiplying another factor.
42pub 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                // SAFETY: backend detection checked all required AVX-512 features.
249                unsafe { ntt_simd::ntt_batch_avx512(a, width) };
250                return;
251            }
252            SimdBackend::Avx2 => {
253                // SAFETY: backend detection checked AVX2.
254                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                // SAFETY: backend detection checked all required AVX-512 features.
330                unsafe { ntt_simd::intt_batch_avx512::<_, false>(a, width) };
331                return;
332            }
333            SimdBackend::Avx2 => {
334                // SAFETY: backend detection checked AVX2.
335                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    // MInt is transparent over u32; lazy residues stay below 2 * MOD and are
353    // normalized before the typed slice is used again.
354    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)] // SIMD intrinsics and raw pointers are confined here
495mod 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        // Limit coefficient growth to leave headroom for FFT roundoff.
848        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        // The 22-bit fallback needs room for three limb products per coefficient.
939        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                            // Block products contain lazy Montgomery residues.
1034                            *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    // Keep limb convolutions below 2^47 at the 2^21 FFT limit.
1116    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    /// Transforms coefficients into the usual NTT frequency order.
1182    fn transform_ntt(t: Self::T, len: usize) -> Self::F {
1183        Self::transform(t, len)
1184    }
1185
1186    /// Inverts a value produced by `transform_ntt`.
1187    fn inverse_transform_ntt(f: Self::F, len: usize) -> Self::T {
1188        Self::inverse_transform(f, len)
1189    }
1190
1191    /// Extends a value produced by `transform_ntt` to twice its length.
1192    /// If `monic`, the input represents a monic degree-`n` polynomial modulo
1193    /// `x^n - 1`, where `n` is the transform length.
1194    fn ntt_doubling(f: Self::F, monic: bool) -> Self::F;
1195
1196    /// Extracts the even coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1197    fn even_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1198
1199    /// Extracts the odd coefficients of `a(x) * b(-x)` in the usual NTT frequency order.
1200    fn odd_mul_normal_neg(f: &Self::F, g: &Self::F) -> Self::F;
1201
1202    /// Multiplies a usual NTT transform by the corresponding prefix of another one.
1203    fn multiply_prefix(f: &mut Self::F, g: &Self::F);
1204
1205    /// Adds the pointwise product of two usual NTT transforms to `sum`.
1206    fn multiply_add(sum: &mut Self::F, f: &Self::F, g: &Self::F);
1207
1208    /// Maximum number of products that can be summed before reconstruction.
1209    /// Both factors must transform canonical coefficients at the supplied transform's length,
1210    /// and each cyclic product must itself be reconstructible.
1211    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            // doubling
2004            {
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            // even_mul_normal_neg
2036            {
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            // odd_mul_normal_neg
2046            {
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            // doubling
2072            {
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            // even_mul_normal_neg
2105            {
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            // odd_mul_normal_neg
2119            {
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}