Skip to main content

add_mod_256

Function add_mod_256 

Source
pub unsafe fn add_mod_256(a: __m256i, b: __m256i, mod_vec: __m256i) -> __m256i
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 158)
146pub unsafe fn add_vec_avx2<M>(
147    a: __m256i,
148    b: __m256i,
149    mod_vec: __m256i,
150    mod2_vec: __m256i,
151) -> __m256i
152where
153    M: Montgomery32NttModulus,
154{
155    if M::MOD < LAZY_THRESHOLD {
156        montgomery_simd::montgomery_add_256(a, b, mod2_vec)
157    } else {
158        montgomery_simd::add_mod_256(a, b, mod_vec)
159    }
160}
161
162pub unsafe fn sub_vec_avx2<M>(
163    a: __m256i,
164    b: __m256i,
165    mod_vec: __m256i,
166    mod2_vec: __m256i,
167) -> __m256i
168where
169    M: Montgomery32NttModulus,
170{
171    if M::MOD < LAZY_THRESHOLD {
172        montgomery_simd::montgomery_sub_256(a, b, mod2_vec)
173    } else {
174        montgomery_simd::sub_mod_256(a, b, mod_vec)
175    }
176}
177
178unsafe fn mul_vec_avx2<M>(a: __m256i, b: __m256i, r_vec: __m256i, mod_vec: __m256i) -> __m256i
179where
180    M: Montgomery32NttModulus,
181{
182    if M::MOD < LAZY_THRESHOLD {
183        montgomery_simd::montgomery_mul_256(a, b, r_vec, mod_vec)
184    } else {
185        montgomery_simd::montgomery_mul_256_canon(a, b, r_vec, mod_vec)
186    }
187}
188
189#[target_feature(enable = "avx2")]
190pub unsafe fn pointwise_multiply_avx2<M>(f: &mut [MInt<M>], g: &[MInt<M>])
191where
192    M: Montgomery32NttModulus,
193{
194    let r_vec = _mm256_set1_epi32(M::R as i32);
195    let mod_vec = _mm256_set1_epi32(M::MOD as i32);
196    let mut i = 0;
197    while i + 8 <= f.len() {
198        let a = _mm256_loadu_si256(f.as_ptr().add(i) as *const __m256i);
199        let b = _mm256_loadu_si256(g.as_ptr().add(i) as *const __m256i);
200        let x = montgomery_simd::montgomery_mul_256_canon(a, b, r_vec, mod_vec);
201        _mm256_storeu_si256(f.as_mut_ptr().add(i) as *mut __m256i, x);
202        i += 8;
203    }
204    while i < f.len() {
205        f[i] *= g[i];
206        i += 1;
207    }
208}
209
210#[target_feature(enable = "avx2")]
211pub unsafe fn pointwise_multiply_add_avx2<M>(sum: &mut [MInt<M>], f: &[MInt<M>], g: &[MInt<M>])
212where
213    M: Montgomery32NttModulus,
214{
215    let r_vec = _mm256_set1_epi32(M::R as i32);
216    let mod_vec = _mm256_set1_epi32(M::MOD as i32);
217    let mut i = 0;
218    while i + 8 <= sum.len() {
219        let s = _mm256_loadu_si256(sum.as_ptr().add(i).cast());
220        let f = _mm256_loadu_si256(f.as_ptr().add(i).cast());
221        let g = _mm256_loadu_si256(g.as_ptr().add(i).cast());
222        let product = montgomery_simd::montgomery_mul_256_canon(f, g, r_vec, mod_vec);
223        _mm256_storeu_si256(
224            sum.as_mut_ptr().add(i).cast(),
225            montgomery_simd::add_mod_256(s, product, mod_vec),
226        );
227        i += 8;
228    }
229    while i < sum.len() {
230        sum[i] += f[i] * g[i];
231        i += 1;
232    }
233}
More examples
Hide additional examples
crates/competitive/src/num/mint/montgomery_dot_product.rs (line 99)
72    pub unsafe fn add_scaled_avx2<M: MontgomeryReduction32>(
73        x: &mut [MInt<M>],
74        y: &[MInt<M>],
75        a: &MInt<M>,
76    ) {
77        // SAFETY: MInt is transparent over u32; preserve the Montgomery representation.
78        let a = *(a as *const MInt<M>).cast::<u32>();
79        let factor = _mm256_set1_epi32(a as i32);
80        let factor_r = _mm256_set1_epi32(a.wrapping_mul(M::R) as i32);
81        let modulus = _mm256_set1_epi32(M::MOD as i32);
82        let end = x.len() / 8 * 8;
83        for i in (0..end).step_by(8) {
84            let value = _mm256_loadu_si256(y.as_ptr().add(i).cast());
85            let odd = _mm256_srli_epi64::<32>(value);
86            let lo = _mm256_mul_epu32(value, factor);
87            let hi = _mm256_mul_epu32(odd, factor);
88            let lo = _mm256_add_epi64(
89                lo,
90                _mm256_mul_epu32(_mm256_mul_epu32(value, factor_r), modulus),
91            );
92            let hi = _mm256_add_epi64(
93                hi,
94                _mm256_mul_epu32(_mm256_mul_epu32(odd, factor_r), modulus),
95            );
96            let product = _mm256_or_si256(_mm256_srli_epi64::<32>(lo), hi);
97            let product = _mm256_min_epu32(product, _mm256_sub_epi32(product, modulus));
98            let old = _mm256_loadu_si256(x.as_ptr().add(i).cast());
99            let sum = add_mod_256(old, product, modulus);
100            _mm256_storeu_si256(x.as_mut_ptr().add(i).cast(), sum);
101        }
102        let a = MInt::new_unchecked(a);
103        for (x, y) in x[end..].iter_mut().zip(&y[end..]) {
104            *x += a * *y;
105        }
106    }