Skip to main content

batch_ntt_simd_backend

Function batch_ntt_simd_backend 

Source
fn batch_ntt_simd_backend(len: usize, width: usize) -> SimdBackend
Examples 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}