unsafe fn dot_soa(
a0: &mut [Complex4],
a1: &mut [Complex4],
b0: &mut [Complex4],
b1: &[Complex4],
)Examples found in repository?
crates/competitive/src/math/mint_fft_convolve.rs (line 271)
257pub unsafe fn convolve_mint_avx2<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
258where
259 M: MIntConvert + MIntConvert<u32>,
260{
261 let len = a.len() + b.len() - 1;
262 let n = len.next_power_of_two() / 2;
263 let modulus = <M as MIntConvert<u32>>::mod_into() as i64;
264 let split = (modulus as f64).sqrt() as i64 + 1;
265 let (mut a0, mut a1) = split_coefficients(a, n, modulus, split);
266 let (mut b0, mut b1) = split_coefficients(b, n, modulus, split);
267 fft_soa(&mut a0);
268 fft_soa(&mut a1);
269 fft_soa(&mut b0);
270 fft_soa(&mut b1);
271 dot_soa(&mut a0, &mut a1, &mut b0, &b1);
272 drop(b1);
273 ifft_soa(&mut a0);
274 ifft_soa(&mut a1);
275 ifft_soa(&mut b0);
276 let split2 = (split * split % modulus) as f64;
277 let split = _mm256_set1_pd(split as f64);
278 let split2 = _mm256_set1_pd(split2);
279 let inverse = _mm256_set1_pd(1.0 / modulus as f64);
280 let modulus = _mm256_set1_pd(modulus as f64);
281 let magic = _mm256_set1_pd((3i64 << 51) as f64);
282 let mut result = vec![MInt::<M>::from(0u32); len];
283 for (block, ((a0, a1), b0)) in a0.iter().zip(&a1).zip(&b0).enumerate() {
284 for (part, (a0, a1, b0)) in [(&a0.re, &a1.re, &b0.re), (&a0.im, &a1.im, &b0.im)]
285 .into_iter()
286 .enumerate()
287 {
288 let a0 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
289 _mm256_load_pd(a0.as_ptr()),
290 );
291 let a1 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
292 _mm256_load_pd(a1.as_ptr()),
293 );
294 let b0 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
295 _mm256_load_pd(b0.as_ptr()),
296 );
297 let a0 = reduce_mod4(a0, modulus, inverse);
298 let a1 = reduce_mod4(a1, modulus, inverse);
299 let b0 = reduce_mod4(b0, modulus, inverse);
300 let value = _mm256_fmadd_pd(b0, split2, _mm256_fmadd_pd(a1, split, a0));
301 let value = reduce_mod4(value, modulus, inverse);
302 let value = _mm256_add_pd(
303 value,
304 _mm256_and_pd(
305 _mm256_cmp_pd::<_CMP_LT_OQ>(value, _mm256_setzero_pd()),
306 modulus,
307 ),
308 );
309 let value = _mm256_sub_epi64(
310 _mm256_castpd_si256(_mm256_add_pd(value, magic)),
311 _mm256_castpd_si256(magic),
312 );
313 let mut lanes = [0i64; 4];
314 _mm256_storeu_si256(lanes.as_mut_ptr().cast(), value);
315 for (lane, value) in lanes.into_iter().enumerate() {
316 let i = block * 4 + lane + part * n;
317 if i < len {
318 let value = value as u32;
319 let modulus = <M as MIntConvert<u32>>::mod_into();
320 // Expose the reduced range to the conversion's remainder operation.
321 result[i] = MInt::<M>::from(if value < modulus {
322 value
323 } else {
324 value % modulus
325 });
326 }
327 }
328 }
329 }
330 result
331}