Skip to main content

montgomery_mul_512_canon

Function montgomery_mul_512_canon 

Source
pub unsafe fn montgomery_mul_512_canon(
    a: __m512i,
    b: __m512i,
    r_vec: __m512i,
    mod_vec: __m512i,
) -> __m512i
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx512.rs (line 51)
44unsafe fn mul_vec_avx512<M>(a: __m512i, b: __m512i, r_vec: __m512i, mod_vec: __m512i) -> __m512i
45where
46    M: Montgomery32NttModulus,
47{
48    if M::MOD < LAZY_THRESHOLD {
49        montgomery_simd::montgomery_mul_512(a, b, r_vec, mod_vec)
50    } else {
51        montgomery_simd::montgomery_mul_512_canon(a, b, r_vec, mod_vec)
52    }
53}
54
55#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
56pub unsafe fn pointwise_multiply_avx512<M>(f: &mut [MInt<M>], g: &[MInt<M>])
57where
58    M: Montgomery32NttModulus,
59{
60    let r_vec = _mm512_set1_epi32(M::R as i32);
61    let mod_vec = _mm512_set1_epi32(M::MOD as i32);
62    let mut i = 0;
63    while i + 16 <= f.len() {
64        let a = _mm512_loadu_si512(f.as_ptr().add(i) as *const __m512i);
65        let b = _mm512_loadu_si512(g.as_ptr().add(i) as *const __m512i);
66        let x = montgomery_simd::montgomery_mul_512_canon(a, b, r_vec, mod_vec);
67        _mm512_storeu_si512(f.as_mut_ptr().add(i) as *mut __m512i, x);
68        i += 16;
69    }
70    while i < f.len() {
71        f[i] *= g[i];
72        i += 1;
73    }
74}
75
76#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
77pub unsafe fn pointwise_multiply_add_avx512<M>(sum: &mut [MInt<M>], f: &[MInt<M>], g: &[MInt<M>])
78where
79    M: Montgomery32NttModulus,
80{
81    let r_vec = _mm512_set1_epi32(M::R as i32);
82    let mod_vec = _mm512_set1_epi32(M::MOD as i32);
83    let mut i = 0;
84    while i + 16 <= sum.len() {
85        let s = _mm512_loadu_si512(sum.as_ptr().add(i).cast());
86        let f = _mm512_loadu_si512(f.as_ptr().add(i).cast());
87        let g = _mm512_loadu_si512(g.as_ptr().add(i).cast());
88        let product = montgomery_simd::montgomery_mul_512_canon(f, g, r_vec, mod_vec);
89        _mm512_storeu_si512(
90            sum.as_mut_ptr().add(i).cast(),
91            montgomery_simd::add_mod_512(s, product, mod_vec),
92        );
93        i += 16;
94    }
95    while i < sum.len() {
96        sum[i] += f[i] * g[i];
97        i += 1;
98    }
99}
100
101#[inline]
102#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
103pub unsafe fn ntt_batch_avx512<M>(a: &mut [MInt<M>], width: usize)
104where
105    M: Montgomery32NttModulus,
106{
107    let n = a.len() / width;
108    if n <= 1 {
109        return;
110    }
111    let ptr = a.as_mut_ptr() as *mut u32;
112    let a = std::slice::from_raw_parts_mut(ptr, a.len());
113    let mod_vec = _mm512_set1_epi32(M::MOD as i32);
114    let mod2_vec = _mm512_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
115    let r_vec = _mm512_set1_epi32(M::R as i32);
116    let imag = M::INFO.root[2];
117    let imag_vec = _mm512_set1_epi32(imag as i32);
118
119    let mut v = n / 2;
120    if n.trailing_zeros() & 1 == 1 {
121        let half = v * width;
122        let mut i = 0;
123        while i + 16 <= half {
124            let x0 = _mm512_loadu_si512(a.as_ptr().add(i) as *const __m512i);
125            let x1 = _mm512_loadu_si512(a.as_ptr().add(half + i) as *const __m512i);
126            let y0 = add_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
127            let y1 = sub_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
128            _mm512_storeu_si512(a.as_mut_ptr().add(i) as *mut __m512i, y0);
129            _mm512_storeu_si512(a.as_mut_ptr().add(half + i) as *mut __m512i, y1);
130            i += 16;
131        }
132        while i < half {
133            let x0 = a[i];
134            let x1 = a[half + i];
135            a[i] = M::mod_add(x0, x1);
136            a[half + i] = M::mod_sub(x0, x1);
137            i += 1;
138        }
139        v >>= 1;
140    }
141    while v > 1 {
142        if width == 1 && v == 2 && a.len() >= 8 && M::MOD < LAZY_THRESHOLD {
143            ntt_avx2::ntt_four_avx2::<M, false>(a);
144            break;
145        }
146        let half = (v >> 1) * width;
147        let mut w1 = M::N1;
148        for (s, block) in a.chunks_exact_mut((v << 1) * width).enumerate() {
149            let base = block.as_mut_ptr();
150            let ll = base;
151            let lr = base.add(half);
152            let rl = base.add(v * width);
153            let rr = base.add(v * width + half);
154            let w2 = M::mod_mul(w1, w1);
155            let w3 = M::mod_mul(w2, w1);
156            let w1v = _mm512_set1_epi32(w1 as i32);
157            let w2v = _mm512_set1_epi32(w2 as i32);
158            let w3v = _mm512_set1_epi32(w3 as i32);
159
160            let mut i = 0;
161            while i + 16 <= half {
162                let x0 = _mm512_loadu_si512(ll.add(i) as *const __m512i);
163                let x1 = _mm512_loadu_si512(lr.add(i) as *const __m512i);
164                let x2 = _mm512_loadu_si512(rl.add(i) as *const __m512i);
165                let x3 = _mm512_loadu_si512(rr.add(i) as *const __m512i);
166
167                let (a1, a2, a3) = if s == 0 {
168                    (x1, x2, x3)
169                } else {
170                    (
171                        mul_vec_avx512::<M>(x1, w1v, r_vec, mod_vec),
172                        mul_vec_avx512::<M>(x2, w2v, r_vec, mod_vec),
173                        mul_vec_avx512::<M>(x3, w3v, r_vec, mod_vec),
174                    )
175                };
176
177                let a0pa2 = add_vec_avx512::<M>(x0, a2, mod_vec, mod2_vec);
178                let a0na2 = sub_vec_avx512::<M>(x0, a2, mod_vec, mod2_vec);
179                let a1pa3 = add_vec_avx512::<M>(a1, a3, mod_vec, mod2_vec);
180                let a1na3 = sub_vec_avx512::<M>(a1, a3, mod_vec, mod2_vec);
181                let a1na3imag = mul_vec_avx512::<M>(a1na3, imag_vec, r_vec, mod_vec);
182
183                let y0 = add_vec_avx512::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
184                let y1 = sub_vec_avx512::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
185                let y2 = add_vec_avx512::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
186                let y3 = sub_vec_avx512::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
187
188                _mm512_storeu_si512(ll.add(i) as *mut __m512i, y0);
189                _mm512_storeu_si512(lr.add(i) as *mut __m512i, y1);
190                _mm512_storeu_si512(rl.add(i) as *mut __m512i, y2);
191                _mm512_storeu_si512(rr.add(i) as *mut __m512i, y3);
192                i += 16;
193            }
194            while i < half {
195                let a0 = normalize_scalar::<M>(*ll.add(i));
196                let a1 = M::mod_mul(normalize_scalar::<M>(*lr.add(i)), w1);
197                let a2 = M::mod_mul(normalize_scalar::<M>(*rl.add(i)), w2);
198                let a3 = M::mod_mul(normalize_scalar::<M>(*rr.add(i)), w3);
199                let a0pa2 = M::mod_add(a0, a2);
200                let a0na2 = M::mod_sub(a0, a2);
201                let a1pa3 = M::mod_add(a1, a3);
202                let a1na3 = M::mod_sub(a1, a3);
203                let a1na3imag = M::mod_mul(a1na3, imag);
204                *ll.add(i) = M::mod_add(a0pa2, a1pa3);
205                *lr.add(i) = M::mod_sub(a0pa2, a1pa3);
206                *rl.add(i) = M::mod_add(a0na2, a1na3imag);
207                *rr.add(i) = M::mod_sub(a0na2, a1na3imag);
208                i += 1;
209            }
210            w1 = M::mod_mul(w1, M::INFO.rate3[s.trailing_ones() as usize]);
211        }
212        v >>= 2;
213    }
214    normalize_avx512::<M>(a);
215}
216
217#[inline]
218#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
219pub unsafe fn intt_batch_avx512<M, const SINGLE: bool>(a: &mut [MInt<M>], width: usize)
220where
221    M: Montgomery32NttModulus,
222{
223    let width = if SINGLE { 1 } else { width };
224    let n = a.len() / width;
225    if n <= 1 {
226        return;
227    }
228    let ptr = a.as_mut_ptr() as *mut u32;
229    let a = std::slice::from_raw_parts_mut(ptr, a.len());
230    let mod_vec = _mm512_set1_epi32(M::MOD as i32);
231    let mod2_vec = _mm512_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
232    let r_vec = _mm512_set1_epi32(M::R as i32);
233    let iimag = M::INFO.inv_root[2];
234    let iimag_vec = _mm512_set1_epi32(iimag as i32);
235
236    let mut v = 1;
237    if width == 1 && a.len() >= 8 && M::MOD < LAZY_THRESHOLD {
238        ntt_avx2::ntt_four_avx2::<M, true>(a);
239        v = 4;
240    }
241    let limit = if n.trailing_zeros() & 1 == 1 {
242        n / 2
243    } else {
244        n
245    };
246    while v < limit {
247        let quarter = v * width;
248        let mut w1 = M::N1;
249        let mut w2 = w1;
250        let mut w3 = w1;
251        for (s, block) in a.chunks_exact_mut((v << 2) * width).enumerate() {
252            let base = block.as_mut_ptr();
253            let ll = base;
254            let lr = base.add(quarter);
255            let rl = base.add(quarter * 2);
256            let rr = base.add(quarter * 3);
257            let w1v = _mm512_set1_epi32(w1 as i32);
258            let w2v = _mm512_set1_epi32(w2 as i32);
259            let w3v = _mm512_set1_epi32(w3 as i32);
260
261            let mut i = 0;
262            while i + 16 <= quarter {
263                let x0 = _mm512_loadu_si512(ll.add(i) as *const __m512i);
264                let x1 = _mm512_loadu_si512(lr.add(i) as *const __m512i);
265                let x2 = _mm512_loadu_si512(rl.add(i) as *const __m512i);
266                let x3 = _mm512_loadu_si512(rr.add(i) as *const __m512i);
267
268                let a0pa1 = add_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
269                let a0na1 = sub_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
270                let a2pa3 = add_vec_avx512::<M>(x2, x3, mod_vec, mod2_vec);
271                let a2na3 = sub_vec_avx512::<M>(x2, x3, mod_vec, mod2_vec);
272                let a2na3iimag = mul_vec_avx512::<M>(a2na3, iimag_vec, r_vec, mod_vec);
273
274                let y0 = add_vec_avx512::<M>(a0pa1, a2pa3, mod_vec, mod2_vec);
275                let y1 = add_vec_avx512::<M>(a0na1, a2na3iimag, mod_vec, mod2_vec);
276                let y2 = sub_vec_avx512::<M>(a0pa1, a2pa3, mod_vec, mod2_vec);
277                let y3 = sub_vec_avx512::<M>(a0na1, a2na3iimag, mod_vec, mod2_vec);
278
279                let (y1, y2, y3) = if s == 0 {
280                    (y1, y2, y3)
281                } else {
282                    (
283                        mul_vec_avx512::<M>(y1, w1v, r_vec, mod_vec),
284                        mul_vec_avx512::<M>(y2, w2v, r_vec, mod_vec),
285                        mul_vec_avx512::<M>(y3, w3v, r_vec, mod_vec),
286                    )
287                };
288
289                _mm512_storeu_si512(ll.add(i) as *mut __m512i, y0);
290                _mm512_storeu_si512(lr.add(i) as *mut __m512i, y1);
291                _mm512_storeu_si512(rl.add(i) as *mut __m512i, y2);
292                _mm512_storeu_si512(rr.add(i) as *mut __m512i, y3);
293                i += 16;
294            }
295            while i < quarter {
296                let a0 = *ll.add(i);
297                let a1 = *lr.add(i);
298                let a2 = *rl.add(i);
299                let a3 = *rr.add(i);
300                let a0pa1 = add_scalar::<M>(a0, a1);
301                let a0na1 = sub_scalar::<M>(a0, a1);
302                let a2pa3 = add_scalar::<M>(a2, a3);
303                let a2na3iimag = mul_scalar::<M>(sub_scalar::<M>(a2, a3), iimag);
304                *ll.add(i) = add_scalar::<M>(a0pa1, a2pa3);
305                *lr.add(i) = mul_scalar::<M>(add_scalar::<M>(a0na1, a2na3iimag), w1);
306                *rl.add(i) = mul_scalar::<M>(sub_scalar::<M>(a0pa1, a2pa3), w2);
307                *rr.add(i) = mul_scalar::<M>(sub_scalar::<M>(a0na1, a2na3iimag), w3);
308                i += 1;
309            }
310            let rate = &M::INFO.inv_rate3_packed[s.trailing_ones() as usize];
311            w1 = M::mod_mul(w1, rate[1]);
312            w2 = M::mod_mul(w2, rate[3]);
313            w3 = M::mod_mul(w3, rate[5]);
314        }
315        v <<= 2;
316    }
317    if n.trailing_zeros() & 1 == 1 {
318        let half = n / 2 * width;
319        let mut i = 0;
320        while i + 16 <= half {
321            let x0 = _mm512_loadu_si512(a.as_ptr().add(i) as *const __m512i);
322            let x1 = _mm512_loadu_si512(a.as_ptr().add(half + i) as *const __m512i);
323            let y0 = add_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
324            let y1 = sub_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
325            _mm512_storeu_si512(a.as_mut_ptr().add(i) as *mut __m512i, y0);
326            _mm512_storeu_si512(a.as_mut_ptr().add(half + i) as *mut __m512i, y1);
327            i += 16;
328        }
329        while i < half {
330            let x0 = a[i];
331            let x1 = a[half + i];
332            a[i] = add_scalar::<M>(x0, x1);
333            a[half + i] = sub_scalar::<M>(x0, x1);
334            i += 1;
335        }
336    }
337    let inv = M::mod_inv(<M as MIntConvert<u32>>::from(n as u32));
338    let inv_vec = _mm512_set1_epi32(inv as i32);
339    let mut i = 0;
340    while i + 16 <= a.len() {
341        let x = _mm512_loadu_si512(a.as_ptr().add(i) as *const __m512i);
342        let y = montgomery_simd::montgomery_mul_512_canon(x, inv_vec, r_vec, mod_vec);
343        _mm512_storeu_si512(a.as_mut_ptr().add(i) as *mut __m512i, y);
344        i += 16;
345    }
346    while i < a.len() {
347        a[i] = M::mod_mul(a[i], inv);
348        i += 1;
349    }
350}