competitive/num/mint/
montgomery_dot_product.rs1use 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 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 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)] mod 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 #[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 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 #[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 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}