Skip to main content

add_mod_512

Function add_mod_512 

Source
pub unsafe fn add_mod_512(a: __m512i, b: __m512i, mod_vec: __m512i) -> __m512i
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx512.rs (line 29)
22unsafe fn add_vec_avx512<M>(a: __m512i, b: __m512i, mod_vec: __m512i, mod2_vec: __m512i) -> __m512i
23where
24    M: Montgomery32NttModulus,
25{
26    if M::MOD < LAZY_THRESHOLD {
27        montgomery_simd::montgomery_add_512(a, b, mod2_vec)
28    } else {
29        montgomery_simd::add_mod_512(a, b, mod_vec)
30    }
31}
32
33unsafe fn sub_vec_avx512<M>(a: __m512i, b: __m512i, mod_vec: __m512i, mod2_vec: __m512i) -> __m512i
34where
35    M: Montgomery32NttModulus,
36{
37    if M::MOD < LAZY_THRESHOLD {
38        montgomery_simd::montgomery_sub_512(a, b, mod2_vec)
39    } else {
40        montgomery_simd::sub_mod_512(a, b, mod_vec)
41    }
42}
43
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}
More examples
Hide additional examples
crates/competitive/src/num/mint/montgomery_dot_product.rs (line 138)
111    pub unsafe fn add_scaled_avx512<M: MontgomeryReduction32>(
112        x: &mut [MInt<M>],
113        y: &[MInt<M>],
114        a: &MInt<M>,
115    ) {
116        // SAFETY: MInt is transparent over u32; preserve the Montgomery representation.
117        let a = *(a as *const MInt<M>).cast::<u32>();
118        let factor = _mm512_set1_epi32(a as i32);
119        let factor_r = _mm512_set1_epi32(a.wrapping_mul(M::R) as i32);
120        let modulus = _mm512_set1_epi32(M::MOD as i32);
121        let end = x.len() / 16 * 16;
122        for i in (0..end).step_by(16) {
123            let value = _mm512_loadu_si512(y.as_ptr().add(i).cast());
124            let odd = _mm512_srli_epi64::<32>(value);
125            let lo = _mm512_mul_epu32(value, factor);
126            let hi = _mm512_mul_epu32(odd, factor);
127            let lo = _mm512_add_epi64(
128                lo,
129                _mm512_mul_epu32(_mm512_mul_epu32(value, factor_r), modulus),
130            );
131            let hi = _mm512_add_epi64(
132                hi,
133                _mm512_mul_epu32(_mm512_mul_epu32(odd, factor_r), modulus),
134            );
135            let product = _mm512_or_si512(_mm512_srli_epi64::<32>(lo), hi);
136            let product = _mm512_min_epu32(product, _mm512_sub_epi32(product, modulus));
137            let old = _mm512_loadu_si512(x.as_ptr().add(i).cast());
138            let sum = add_mod_512(old, product, modulus);
139            _mm512_storeu_si512(x.as_mut_ptr().add(i).cast(), sum);
140        }
141        let a = MInt::new_unchecked(a);
142        for (x, y) in x[end..].iter_mut().zip(&y[end..]) {
143            *x += a * *y;
144        }
145    }