pub unsafe fn add_mod_256(a: __m256i, b: __m256i, mod_vec: __m256i) -> __m256iExamples 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
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 }