Skip to main content

mod_mul

Function mod_mul 

Source
const fn mod_mul(x: u32, y: u32, p: u32, r: u32) -> u32
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 71)
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    }
More examples
Hide additional examples
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 42)
23pub unsafe fn ntt_four_avx2<M, const INVERSE: bool>(a: &mut [u32])
24where
25    M: Montgomery32NttModulus,
26{
27    let roots = if INVERSE {
28        &M::INFO.inv_root
29    } else {
30        &M::INFO.root
31    };
32    let rates = &const {
33        let info = M::INFO;
34        let (roots, rates) = if INVERSE {
35            (info.inv_root, info.inv_rate3_packed)
36        } else {
37            (info.root, info.rate3_packed)
38        };
39        let mut result = [[0u32; 4]; 32];
40        let mut i = 0;
41        while i < 32 {
42            let r = mod_mul(rates[i][1], roots[3], M::MOD, M::R);
43            let r2 = mod_mul(r, r, M::MOD, M::R);
44            result[i] = [M::N1, r, r2, mod_mul(r2, r, M::MOD, M::R)];
45            i += 1;
46        }
47        result
48    };
49    let one = _mm256_set1_epi32(M::N1 as i32);
50    let modulus = _mm256_set1_epi32(M::MOD as i32);
51    let modulus2 = _mm256_set1_epi32((M::MOD * 2) as i32);
52    let r = _mm256_set1_epi32(M::R as i32);
53    let root3 = M::mod_mul(roots[2], roots[3]) as i32;
54    let step = _mm256_setr_epi32(
55        M::N1 as i32,
56        roots[3] as i32,
57        roots[2] as i32,
58        root3,
59        M::N1 as i32,
60        roots[3] as i32,
61        roots[2] as i32,
62        root3,
63    );
64    let mut twiddle = _mm256_blend_epi32::<0xf0>(one, step);
65    let imag = if INVERSE {
66        _mm256_setr_epi32(
67            M::N1 as i32,
68            roots[2] as i32,
69            M::N1 as i32,
70            roots[2] as i32,
71            M::N1 as i32,
72            roots[2] as i32,
73            M::N1 as i32,
74            roots[2] as i32,
75        )
76    } else {
77        _mm256_setr_epi32(
78            M::N1 as i32,
79            M::N1 as i32,
80            roots[2] as i32,
81            roots[2] as i32,
82            M::N1 as i32,
83            M::N1 as i32,
84            roots[2] as i32,
85            roots[2] as i32,
86        )
87    };
88    for (s, a) in a.as_chunks_mut::<8>().0.iter_mut().enumerate() {
89        let mut x = _mm256_loadu_si256(a.as_ptr().cast());
90        if !INVERSE {
91            x = montgomery_simd::montgomery_mul_256(x, twiddle, r, modulus);
92        }
93        let pair = if INVERSE {
94            let y = _mm256_shuffle_epi32::<0xb1>(x);
95            let sum = montgomery_simd::montgomery_add_256(x, y, modulus2);
96            let diff = montgomery_simd::montgomery_sub_256(x, y, modulus2);
97            let sum = _mm256_shuffle_epi32::<0x88>(sum);
98            let diff = _mm256_shuffle_epi32::<0x88>(diff);
99            let diff = montgomery_simd::montgomery_mul_256(diff, imag, r, modulus);
100            _mm256_unpacklo_epi64(sum, diff)
101        } else {
102            let y = _mm256_shuffle_epi32::<0x4e>(x);
103            let sum = montgomery_simd::montgomery_add_256(x, y, modulus2);
104            let diff = montgomery_simd::montgomery_sub_256(x, y, modulus2);
105            _mm256_unpacklo_epi64(sum, diff)
106        };
107        let left = _mm256_shuffle_epi32::<0xa0>(pair);
108        let mut right = _mm256_shuffle_epi32::<0xf5>(pair);
109        if !INVERSE {
110            right = montgomery_simd::montgomery_mul_256(right, imag, r, modulus);
111        }
112        let sum = montgomery_simd::montgomery_add_256(left, right, modulus2);
113        let diff = montgomery_simd::montgomery_sub_256(left, right, modulus2);
114        let mut value = _mm256_blend_epi32::<0xaa>(sum, diff);
115        if INVERSE {
116            value = _mm256_shuffle_epi32::<0xd8>(value);
117            value = montgomery_simd::montgomery_mul_256(value, twiddle, r, modulus);
118        }
119        _mm256_storeu_si256(a.as_mut_ptr().cast(), value);
120        let rate = _mm256_broadcastsi128_si256(_mm_loadu_si128(
121            rates[s.trailing_ones() as usize + 1].as_ptr().cast(),
122        ));
123        twiddle = montgomery_simd::montgomery_mul_256(twiddle, rate, r, modulus);
124    }
125}