fn mul_scalar<M>(x: u32, y: u32) -> u32where
M: Montgomery32NttModulus,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
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}