Skip to main content

mul_scalar

Function mul_scalar 

Source
fn mul_scalar<M>(x: u32, y: u32) -> u32
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 378)
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}
More examples
Hide additional examples
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 333)
237pub unsafe fn ntt_batch_avx2<M, const PARTIAL: bool>(a: &mut [MInt<M>], width: usize)
238where
239    M: Montgomery32NttModulus,
240{
241    let n = a.len() / width;
242    if n <= 1 {
243        return;
244    }
245    let ptr = a.as_mut_ptr() as *mut u32;
246    let a = std::slice::from_raw_parts_mut(ptr, a.len());
247    let mod_vec = _mm256_set1_epi32(M::MOD as i32);
248    let mod2_vec = _mm256_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
249    let r_vec = _mm256_set1_epi32(M::R as i32);
250    let imag = M::INFO.root[2];
251    let imag_vec = _mm256_set1_epi32(imag as i32);
252
253    let mut v = n / 2;
254    if n.trailing_zeros() & 1 == 1 {
255        let half = v * width;
256        let step = if PARTIAL && half == 4 { 4 } else { 8 };
257        let mut i = 0;
258        while i + step <= half {
259            let x0 = load_ntt_avx2(a.as_ptr().add(i), step);
260            let x1 = load_ntt_avx2(a.as_ptr().add(half + i), step);
261            let y0 = add_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
262            let y1 = sub_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
263            store_ntt_avx2(a.as_mut_ptr().add(i), y0, step);
264            store_ntt_avx2(a.as_mut_ptr().add(half + i), y1, step);
265            i += step;
266        }
267        while i < half {
268            let x0 = a[i];
269            let x1 = a[half + i];
270            a[i] = add_scalar::<M>(x0, x1);
271            a[half + i] = sub_scalar::<M>(x0, x1);
272            i += 1;
273        }
274        v >>= 1;
275    }
276    while v > 1 {
277        if width == 1 && v == 2 && a.len() >= 8 && M::MOD < LAZY_THRESHOLD {
278            ntt_four_avx2::<M, false>(a);
279            break;
280        }
281        let half = (v >> 1) * width;
282        let step = if PARTIAL && half == 4 { 4 } else { 8 };
283        let mut w1 = M::N1;
284        let mut w2 = w1;
285        let mut w3 = w1;
286        for (s, block) in a.chunks_exact_mut((v << 1) * width).enumerate() {
287            let base = block.as_mut_ptr();
288            let ll = base;
289            let lr = base.add(half);
290            let rl = base.add(v * width);
291            let rr = base.add(v * width + half);
292
293            let w1v = _mm256_set1_epi32(w1 as i32);
294            let w2v = _mm256_set1_epi32(w2 as i32);
295            let w3v = _mm256_set1_epi32(w3 as i32);
296
297            let mut i = 0;
298            while i + step <= half {
299                let x0 = load_ntt_avx2(ll.add(i), step);
300                let x1 = load_ntt_avx2(lr.add(i), step);
301                let x2 = load_ntt_avx2(rl.add(i), step);
302                let x3 = load_ntt_avx2(rr.add(i), step);
303
304                let (a1, a2, a3) = if s == 0 {
305                    (x1, x2, x3)
306                } else {
307                    (
308                        mul_vec_avx2::<M>(x1, w1v, r_vec, mod_vec),
309                        mul_vec_avx2::<M>(x2, w2v, r_vec, mod_vec),
310                        mul_vec_avx2::<M>(x3, w3v, r_vec, mod_vec),
311                    )
312                };
313
314                let a0pa2 = add_vec_avx2::<M>(x0, a2, mod_vec, mod2_vec);
315                let a0na2 = sub_vec_avx2::<M>(x0, a2, mod_vec, mod2_vec);
316                let a1pa3 = add_vec_avx2::<M>(a1, a3, mod_vec, mod2_vec);
317                let a1na3 = sub_vec_avx2::<M>(a1, a3, mod_vec, mod2_vec);
318                let a1na3imag = mul_vec_avx2::<M>(a1na3, imag_vec, r_vec, mod_vec);
319
320                let y0 = add_vec_avx2::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
321                let y1 = sub_vec_avx2::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
322                let y2 = add_vec_avx2::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
323                let y3 = sub_vec_avx2::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
324
325                store_ntt_avx2(ll.add(i), y0, step);
326                store_ntt_avx2(lr.add(i), y1, step);
327                store_ntt_avx2(rl.add(i), y2, step);
328                store_ntt_avx2(rr.add(i), y3, step);
329                i += step;
330            }
331            while i < half {
332                let a0 = *ll.add(i);
333                let a1 = mul_scalar::<M>(*lr.add(i), w1);
334                let a2 = mul_scalar::<M>(*rl.add(i), w2);
335                let a3 = mul_scalar::<M>(*rr.add(i), w3);
336                let a0pa2 = add_scalar::<M>(a0, a2);
337                let a0na2 = sub_scalar::<M>(a0, a2);
338                let a1pa3 = add_scalar::<M>(a1, a3);
339                let a1na3 = sub_scalar::<M>(a1, a3);
340                let a1na3imag = mul_scalar::<M>(a1na3, imag);
341                *ll.add(i) = add_scalar::<M>(a0pa2, a1pa3);
342                *lr.add(i) = sub_scalar::<M>(a0pa2, a1pa3);
343                *rl.add(i) = add_scalar::<M>(a0na2, a1na3imag);
344                *rr.add(i) = sub_scalar::<M>(a0na2, a1na3imag);
345                i += 1;
346            }
347            let rate = &M::INFO.rate3_packed[s.trailing_ones() as usize];
348            w1 = M::mod_mul(w1, rate[1]);
349            w2 = M::mod_mul(w2, rate[3]);
350            w3 = M::mod_mul(w3, rate[5]);
351        }
352        v >>= 2;
353    }
354    normalize_avx2::<M>(a);
355}
356
357#[inline]
358#[target_feature(enable = "avx2")]
359pub unsafe fn intt_batch_avx2<M, const PARTIAL: bool>(a: &mut [MInt<M>], width: usize)
360where
361    M: Montgomery32NttModulus,
362{
363    let n = a.len() / width;
364    if n <= 1 {
365        return;
366    }
367    let ptr = a.as_mut_ptr() as *mut u32;
368    let a = std::slice::from_raw_parts_mut(ptr, a.len());
369    let mod_vec = _mm256_set1_epi32(M::MOD as i32);
370    let mod2_vec = _mm256_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
371    let r_vec = _mm256_set1_epi32(M::R as i32);
372    let iimag = M::INFO.inv_root[2];
373    let iimag_vec = _mm256_set1_epi32(iimag as i32);
374
375    let mut v = 1;
376    if width == 1 && a.len() >= 8 && M::MOD < LAZY_THRESHOLD {
377        ntt_four_avx2::<M, true>(a);
378        v = 4;
379    }
380    let limit = if n.trailing_zeros() & 1 == 1 {
381        n / 2
382    } else {
383        n
384    };
385    while v < limit {
386        let quarter = v * width;
387        let step = if PARTIAL && quarter == 4 { 4 } else { 8 };
388        let mut w1 = M::N1;
389        let mut w2 = w1;
390        let mut w3 = w1;
391        for (s, block) in a.chunks_exact_mut((v << 2) * width).enumerate() {
392            let base = block.as_mut_ptr();
393            let ll = base;
394            let lr = base.add(quarter);
395            let rl = base.add(quarter * 2);
396            let rr = base.add(quarter * 3);
397
398            let w1v = _mm256_set1_epi32(w1 as i32);
399            let w2v = _mm256_set1_epi32(w2 as i32);
400            let w3v = _mm256_set1_epi32(w3 as i32);
401
402            let mut i = 0;
403            while i + step <= quarter {
404                let x0 = load_ntt_avx2(ll.add(i), step);
405                let x1 = load_ntt_avx2(lr.add(i), step);
406                let x2 = load_ntt_avx2(rl.add(i), step);
407                let x3 = load_ntt_avx2(rr.add(i), step);
408
409                let a0pa1 = add_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
410                let a0na1 = sub_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
411                let a2pa3 = add_vec_avx2::<M>(x2, x3, mod_vec, mod2_vec);
412                let a2na3 = sub_vec_avx2::<M>(x2, x3, mod_vec, mod2_vec);
413                let a2na3iimag = mul_vec_avx2::<M>(a2na3, iimag_vec, r_vec, mod_vec);
414
415                let y0 = add_vec_avx2::<M>(a0pa1, a2pa3, mod_vec, mod2_vec);
416                let y1 = add_vec_avx2::<M>(a0na1, a2na3iimag, mod_vec, mod2_vec);
417                let y2 = sub_vec_avx2::<M>(a0pa1, a2pa3, mod_vec, mod2_vec);
418                let y3 = sub_vec_avx2::<M>(a0na1, a2na3iimag, mod_vec, mod2_vec);
419
420                let (y1, y2, y3) = if s == 0 {
421                    (y1, y2, y3)
422                } else {
423                    (
424                        mul_vec_avx2::<M>(y1, w1v, r_vec, mod_vec),
425                        mul_vec_avx2::<M>(y2, w2v, r_vec, mod_vec),
426                        mul_vec_avx2::<M>(y3, w3v, r_vec, mod_vec),
427                    )
428                };
429
430                store_ntt_avx2(ll.add(i), y0, step);
431                store_ntt_avx2(lr.add(i), y1, step);
432                store_ntt_avx2(rl.add(i), y2, step);
433                store_ntt_avx2(rr.add(i), y3, step);
434                i += step;
435            }
436            while i < quarter {
437                let a0 = *ll.add(i);
438                let a1 = *lr.add(i);
439                let a2 = *rl.add(i);
440                let a3 = *rr.add(i);
441                let a0pa1 = add_scalar::<M>(a0, a1);
442                let a0na1 = sub_scalar::<M>(a0, a1);
443                let a2pa3 = add_scalar::<M>(a2, a3);
444                let a2na3iimag = mul_scalar::<M>(sub_scalar::<M>(a2, a3), iimag);
445                *ll.add(i) = add_scalar::<M>(a0pa1, a2pa3);
446                *lr.add(i) = mul_scalar::<M>(add_scalar::<M>(a0na1, a2na3iimag), w1);
447                *rl.add(i) = mul_scalar::<M>(sub_scalar::<M>(a0pa1, a2pa3), w2);
448                *rr.add(i) = mul_scalar::<M>(sub_scalar::<M>(a0na1, a2na3iimag), w3);
449                i += 1;
450            }
451            let rate = &M::INFO.inv_rate3_packed[s.trailing_ones() as usize];
452            w1 = M::mod_mul(w1, rate[1]);
453            w2 = M::mod_mul(w2, rate[3]);
454            w3 = M::mod_mul(w3, rate[5]);
455        }
456        v <<= 2;
457    }
458    if n.trailing_zeros() & 1 == 1 {
459        let half = (n >> 1) * width;
460        let step = if PARTIAL && half == 4 { 4 } else { 8 };
461        let mut i = 0;
462        while i + step <= half {
463            let x0 = load_ntt_avx2(a.as_ptr().add(i), step);
464            let x1 = load_ntt_avx2(a.as_ptr().add(half + i), step);
465            let y0 = add_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
466            let y1 = sub_vec_avx2::<M>(x0, x1, mod_vec, mod2_vec);
467            store_ntt_avx2(a.as_mut_ptr().add(i), y0, step);
468            store_ntt_avx2(a.as_mut_ptr().add(half + i), y1, step);
469            i += step;
470        }
471        while i < half {
472            let x0 = a[i];
473            let x1 = a[half + i];
474            a[i] = add_scalar::<M>(x0, x1);
475            a[half + i] = sub_scalar::<M>(x0, x1);
476            i += 1;
477        }
478    }
479    let inv = M::mod_inv(<M as MIntConvert<u32>>::from(n as u32));
480    let inv_vec = _mm256_set1_epi32(inv as i32);
481    let mut i = 0;
482    while i + 8 <= a.len() {
483        let x = _mm256_loadu_si256(a.as_ptr().add(i) as *const __m256i);
484        let y = montgomery_simd::montgomery_mul_256_canon(x, inv_vec, r_vec, mod_vec);
485        _mm256_storeu_si256(a.as_mut_ptr().add(i) as *mut __m256i, y);
486        i += 8;
487    }
488    while i < a.len() {
489        a[i] = M::mod_mul(a[i], inv);
490        i += 1;
491    }
492}
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx512.rs (line 303)
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}