fn batch_ntt_simd_backend(len: usize, width: usize) -> SimdBackendExamples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 246)
240fn ntt_batch<M>(a: &mut [MInt<M>], width: usize)
241where
242 M: Montgomery32NttModulus,
243{
244 #[cfg(target_arch = "x86_64")]
245 {
246 match batch_ntt_simd_backend(a.len(), width) {
247 SimdBackend::Avx512 => {
248 // SAFETY: backend detection checked all required AVX-512 features.
249 unsafe { ntt_simd::ntt_batch_avx512(a, width) };
250 return;
251 }
252 SimdBackend::Avx2 => {
253 // SAFETY: backend detection checked AVX2.
254 unsafe { ntt_simd::ntt_batch_avx2::<_, false>(a, width) };
255 return;
256 }
257 SimdBackend::Scalar => {}
258 }
259 }
260 ntt_batch_scalar(a, width);
261}
262
263fn ntt_batch_scalar<M>(a: &mut [MInt<M>], width: usize)
264where
265 M: Montgomery32NttModulus,
266{
267 let n = a.len() / width;
268 if n <= 1 {
269 return;
270 }
271 let mut v = n / 2;
272 if n.trailing_zeros() & 1 == 1 {
273 let (l, r) = a.split_at_mut(v * width);
274 for (x0, x1) in l.iter_mut().zip(r) {
275 let a0 = *x0;
276 let a1 = *x1;
277 *x0 = a0 + a1;
278 *x1 = a0 - a1;
279 }
280 v >>= 1;
281 }
282 let imag = MInt::<M>::new_unchecked(M::INFO.root[2]);
283 while v > 1 {
284 let mut w1 = MInt::<M>::one();
285 let mut w2 = w1;
286 let mut w3 = w1;
287 for (s, a) in a.chunks_exact_mut((v << 1) * width).enumerate() {
288 let (l, r) = a.split_at_mut(v * width);
289 let (ll, lr) = l.split_at_mut((v >> 1) * width);
290 let (rl, rr) = r.split_at_mut((v >> 1) * width);
291 for (((x0, x1), x2), x3) in ll.iter_mut().zip(lr).zip(rl).zip(rr) {
292 let a0 = *x0;
293 let a1 = *x1 * w1;
294 let a2 = *x2 * w2;
295 let a3 = *x3 * w3;
296 let a0pa2 = a0 + a2;
297 let a0na2 = a0 - a2;
298 let a1pa3 = a1 + a3;
299 let a1na3imag = (a1 - a3) * imag;
300 *x0 = a0pa2 + a1pa3;
301 *x1 = a0pa2 - a1pa3;
302 *x2 = a0na2 + a1na3imag;
303 *x3 = a0na2 - a1na3imag;
304 }
305 let rate = &M::INFO.rate3_packed[s.trailing_ones() as usize];
306 w1 *= MInt::<M>::new_unchecked(rate[1]);
307 w2 *= MInt::<M>::new_unchecked(rate[3]);
308 w3 *= MInt::<M>::new_unchecked(rate[5]);
309 }
310 v >>= 2;
311 }
312}
313
314fn intt_scalar<M>(a: &mut [MInt<M>])
315where
316 M: Montgomery32NttModulus,
317{
318 intt_batch_scalar(a, 1);
319}
320
321fn intt_batch<M>(a: &mut [MInt<M>], width: usize)
322where
323 M: Montgomery32NttModulus,
324{
325 #[cfg(target_arch = "x86_64")]
326 {
327 match batch_ntt_simd_backend(a.len(), width) {
328 SimdBackend::Avx512 => {
329 // SAFETY: backend detection checked all required AVX-512 features.
330 unsafe { ntt_simd::intt_batch_avx512::<_, false>(a, width) };
331 return;
332 }
333 SimdBackend::Avx2 => {
334 // SAFETY: backend detection checked AVX2.
335 unsafe { ntt_simd::intt_batch_avx2::<_, false>(a, width) };
336 return;
337 }
338 SimdBackend::Scalar => {}
339 }
340 }
341 intt_batch_scalar(a, width);
342}