Skip to main content

competitive/num/mint/
montgomery_dot_product.rs

1use super::{MInt, MIntBase, MIntDotProduct, montgomery::MontgomeryReduction32};
2#[cfg(target_arch = "x86_64")]
3use super::{avx512_enabled, avx512_supported};
4
5impl<M> MIntDotProduct for M
6where
7    M: MontgomeryReduction32,
8{
9    fn try_matrix_product(
10        _a: &[Vec<MInt<Self>>],
11        _b: &[Vec<MInt<Self>>],
12    ) -> Option<Vec<Vec<MInt<Self>>>> {
13        #[cfg(target_arch = "x86_64")]
14        if _a.len() >= 32
15            && _b.len() >= 32
16            && _b[0].len() >= 32
17            && <Self as MontgomeryReduction32>::MOD > 1
18            && <Self as MontgomeryReduction32>::MOD < 1 << 30
19            && <Self as MontgomeryReduction32>::MOD % 2 == 1
20            && is_x86_feature_detected!("avx2")
21        {
22            return Some(unsafe { MInt::matrix_product_avx2(_a, _b, 1) });
23        }
24        None
25    }
26    fn dot_product(x: &[MInt<Self>], y: &[MInt<Self>]) -> MInt<Self> {
27        assert_eq!(x.len(), y.len());
28        // reduce() needs sum < modulus * 2^32 to return a canonical residue.
29        let modulus = <Self as MontgomeryReduction32>::MOD as u64;
30        let block = (((modulus << 32) - 1) / ((modulus - 1) * (modulus - 1))).min(16) as usize;
31        let mut result = 0;
32        for (x, y) in x.chunks(block).zip(y.chunks(block)) {
33            // SAFETY: MInt is transparent over u32 and both slices have equal lengths.
34            let sum = unsafe {
35                let a = std::slice::from_raw_parts(x.as_ptr().cast::<u32>(), x.len());
36                let b = std::slice::from_raw_parts(y.as_ptr().cast::<u32>(), y.len());
37                a.iter().zip(b).map(|(&a, &b)| a as u64 * b as u64).sum()
38            };
39            result = Self::mod_add(result, Self::reduce(sum));
40        }
41        MInt::new_unchecked(result)
42    }
43    fn add_scaled_assign(x: &mut [MInt<Self>], y: &[MInt<Self>], a: &MInt<Self>) {
44        assert_eq!(x.len(), y.len());
45        #[cfg(target_arch = "x86_64")]
46        if x.len() >= 16 {
47            if x.len() >= 64 && avx512_enabled() && avx512_supported() {
48                unsafe { simd::add_scaled_avx512::<Self>(x, y, a) };
49                return;
50            }
51            if is_x86_feature_detected!("avx2") {
52                unsafe { simd::add_scaled_avx2::<Self>(x, y, a) };
53                return;
54            }
55        }
56        for (x, y) in x.iter_mut().zip(y) {
57            *x += *a * *y;
58        }
59    }
60}
61
62#[cfg(target_arch = "x86_64")]
63#[allow(unsafe_op_in_unsafe_fn)] // SIMD intrinsics and raw pointers are confined here
64mod simd {
65    use super::super::montgomery_simd::{add_mod_256, add_mod_512};
66    use super::{MInt, MontgomeryReduction32};
67    use std::arch::x86_64::*;
68
69    /// # Safety
70    /// AVX2 must be available, and `x` and `y` must have equal lengths.
71    #[target_feature(enable = "avx2")]
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    }
107
108    /// # Safety
109    /// AVX-512F/DQ/CD/BW/VL must be available, and `x` and `y` must have equal lengths.
110    #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
111    pub unsafe fn add_scaled_avx512<M: MontgomeryReduction32>(
112        x: &mut [MInt<M>],
113        y: &[MInt<M>],
114        a: &MInt<M>,
115    ) {
116        // SAFETY: MInt is transparent over u32; preserve the Montgomery representation.
117        let a = *(a as *const MInt<M>).cast::<u32>();
118        let factor = _mm512_set1_epi32(a as i32);
119        let factor_r = _mm512_set1_epi32(a.wrapping_mul(M::R) as i32);
120        let modulus = _mm512_set1_epi32(M::MOD as i32);
121        let end = x.len() / 16 * 16;
122        for i in (0..end).step_by(16) {
123            let value = _mm512_loadu_si512(y.as_ptr().add(i).cast());
124            let odd = _mm512_srli_epi64::<32>(value);
125            let lo = _mm512_mul_epu32(value, factor);
126            let hi = _mm512_mul_epu32(odd, factor);
127            let lo = _mm512_add_epi64(
128                lo,
129                _mm512_mul_epu32(_mm512_mul_epu32(value, factor_r), modulus),
130            );
131            let hi = _mm512_add_epi64(
132                hi,
133                _mm512_mul_epu32(_mm512_mul_epu32(odd, factor_r), modulus),
134            );
135            let product = _mm512_or_si512(_mm512_srli_epi64::<32>(lo), hi);
136            let product = _mm512_min_epu32(product, _mm512_sub_epi32(product, modulus));
137            let old = _mm512_loadu_si512(x.as_ptr().add(i).cast());
138            let sum = add_mod_512(old, product, modulus);
139            _mm512_storeu_si512(x.as_mut_ptr().add(i).cast(), sum);
140        }
141        let a = MInt::new_unchecked(a);
142        for (x, y) in x[end..].iter_mut().zip(&y[end..]) {
143            *x += a * *y;
144        }
145    }
146}