1#[cfg(target_arch = "x86_64")]
2use super::avx512_enabled;
3use super::{MInt, MIntBase, MIntDotProduct, mint_basic::*};
4
5#[macro_export]
6macro_rules! impl_basic_mint_dot_product {
7 (u32, u64; $($name:ident),* $(,)?) => {
8 $($crate::impl_basic_mint_dot_product!(@impl #[inline(always)] simd32, $name, u32, u64);)*
9 };
10 ($basety:ty, $upperty:ty; $($name:ident),* $(,)?) => {
11 $($crate::impl_basic_mint_dot_product!(@impl #[inline] scalar, $name, $basety, $upperty);)*
12 };
13 (@impl #[$inline:meta] $kind:ident, $name:ident, $basety:ty, $upperty:ty) => {
14 impl MIntDotProduct for $name {
15 fn try_matrix_product(_a: &[Vec<MInt<Self>>], _b: &[Vec<MInt<Self>>]) -> Option<Vec<Vec<MInt<Self>>>> {
16 $crate::impl_basic_mint_dot_product!(@matrix_product $kind, _a, _b)
17 }
18 #[$inline]
19 fn dot_product(x: &[MInt<Self>], y: &[MInt<Self>]) -> MInt<Self> {
20 $crate::impl_basic_mint_dot_product!(@dot_product $kind, $name, x, y, $basety, $upperty)
21 }
22 #[inline]
23 fn add_scaled_assign(x: &mut [MInt<Self>], y: &[MInt<Self>], a: &MInt<Self>) {
24 assert_eq!(x.len(), y.len());
25 $crate::impl_basic_mint_dot_product!(@add_scaled $kind, $name, x, y, a);
26 for (x, y) in x.iter_mut().zip(y) { *x += *a * *y; }
27 }
28 }
29 $crate::impl_basic_mint_dot_product!(@simd_functions $kind, $name);
30 };
31 (@dot_product scalar, $name:ident, $x:ident, $y:ident, $basety:ty, $upperty:ty) => {{
32 assert_eq!($x.len(), $y.len());
33 let modulus = Self::get_mod() as $upperty;
34 let max_value = modulus - 1;
35 let block = ((<$upperty>::MAX - max_value) / (max_value * max_value).max(1)).min(64) as usize;
36 let mut result = 0 as $upperty;
37 for (x, y) in $x.chunks(block).zip($y.chunks(block)) {
38 let sum: $upperty = x
39 .iter()
40 .zip(y)
41 .map(|(&x, &y)| x.inner() as $upperty * y.inner() as $upperty)
42 .sum();
43 result += sum % modulus;
44 if result >= modulus {
45 result -= modulus;
46 }
47 }
48 MInt::new_unchecked(result as $basety)
49 }};
50 (@dot_product simd32, $name:ident, $x:ident, $y:ident, $basety:ty, $upperty:ty) => {{
51 #[cfg(target_arch = "x86_64")]
52 {
53 if $x.len() >= 32 {
54 if $x.len() >= 512
55 && avx512_enabled()
56 && is_x86_feature_detected!("avx512f")
57 {
58 return MInt::new_unchecked(unsafe { $name::dot_product_avx512($x, $y) });
59 }
60 if is_x86_feature_detected!("avx2") {
61 return MInt::new_unchecked(unsafe { $name::dot_product_avx2($x, $y) });
62 }
63 }
64 }
65 $crate::impl_basic_mint_dot_product!(@dot_product scalar, $name, $x, $y, $basety, $upperty)
66 }};
67 (@matrix_product scalar, $a:ident, $b:ident) => { None };
68 (@matrix_product simd32, $a:ident, $b:ident) => {{
69 #[cfg(target_arch = "x86_64")]
70 if $a.len() >= 32 && $b.len() >= 32 && $b[0].len() >= 32
71 && Self::get_mod() > 1 && Self::get_mod() < 1 << 30
72 && Self::get_mod() % 2 == 1 && is_x86_feature_detected!("avx2")
73 {
74 let scale = ((1u64 << 32) % Self::get_mod() as u64) as u32;
75 return Some(unsafe { MInt::matrix_product_avx2($a, $b, scale) });
76 }
77 None
78 }};
79 (@add_scaled scalar, $name:ident, $x:ident, $y:ident, $a:ident) => {};
80 (@add_scaled simd32, $name:ident, $x:ident, $y:ident, $a:ident) => {
81 #[cfg(target_arch = "x86_64")]
82 if $x.len() >= 16 && Self::get_mod() <= 1 << 31 {
83 if $x.len() >= 64 && avx512_enabled() && is_x86_feature_detected!("avx512f") {
84 unsafe { Self::add_scaled_avx512($x, $y, $a.inner()) };
85 return;
86 }
87 if is_x86_feature_detected!("avx2") {
88 unsafe { Self::add_scaled_avx2($x, $y, $a.inner()) };
89 return;
90 }
91 }
92 };
93 (@simd_functions scalar, $name:ident) => {};
94 (@simd_functions simd32, $name:ident) => {
95 #[cfg(target_arch = "x86_64")]
96 impl $name {
97 #[allow(unsafe_op_in_unsafe_fn)]
98 #[target_feature(enable = "avx2")]
99 unsafe fn add_scaled_avx2(x: &mut [MInt<Self>], y: &[MInt<Self>], a: u32) {
100 use std::arch::x86_64::*;
101 let modulus = _mm256_set1_epi32(Self::get_mod() as i32);
102 let factor = _mm256_set1_epi32(a as i32);
103 let quotient = _mm256_set1_epi32((((a as u64) << 32) / Self::get_mod() as u64) as i32);
105 let update = |old, value| {
106 let lo = _mm256_srli_epi64::<32>(_mm256_mul_epu32(value, quotient));
107 let hi = _mm256_slli_epi64::<32>(_mm256_srli_epi64::<32>(_mm256_mul_epu32(_mm256_srli_epi64::<32>(value), quotient)));
108 let q = _mm256_or_si256(lo, hi);
109 let product = _mm256_sub_epi32(_mm256_mullo_epi32(value, factor), _mm256_mullo_epi32(q, modulus));
110 let product = _mm256_min_epu32(product, _mm256_sub_epi32(product, modulus));
111 let sum = _mm256_add_epi32(old, product);
112 _mm256_min_epu32(sum, _mm256_sub_epi32(sum, modulus))
113 };
114 let end = x.len() / 8 * 8;
115 for i in (0..end).step_by(8) {
116 let value = _mm256_loadu_si256(y.as_ptr().add(i).cast());
117 let old = _mm256_loadu_si256(x.as_ptr().add(i).cast());
118 _mm256_storeu_si256(x.as_mut_ptr().add(i).cast(), update(old, value));
119 }
120 if end < x.len() {
121 let mask = _mm256_cmpgt_epi32(_mm256_set1_epi32((x.len() - end) as i32), _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
122 let value = _mm256_maskload_epi32(y.as_ptr().add(end).cast(), mask);
123 let old = _mm256_maskload_epi32(x.as_ptr().add(end).cast(), mask);
124 _mm256_maskstore_epi32(x.as_mut_ptr().add(end).cast(), mask, update(old, value));
125 }
126 }
127
128 #[allow(unsafe_op_in_unsafe_fn)]
129 #[target_feature(enable = "avx512f")]
130 unsafe fn add_scaled_avx512(x: &mut [MInt<Self>], y: &[MInt<Self>], a: u32) {
131 use std::arch::x86_64::*;
132 let modulus = _mm512_set1_epi32(Self::get_mod() as i32);
133 let factor = _mm512_set1_epi32(a as i32);
134 let quotient = _mm512_set1_epi32((((a as u64) << 32) / Self::get_mod() as u64) as i32);
136 let end = x.len() / 16 * 16;
137 for i in (0..end).step_by(16) {
138 let value = _mm512_loadu_si512(y.as_ptr().add(i).cast());
139 let lo = _mm512_srli_epi64::<32>(_mm512_mul_epu32(value, quotient));
140 let hi = _mm512_slli_epi64::<32>(_mm512_srli_epi64::<32>(_mm512_mul_epu32(_mm512_srli_epi64::<32>(value), quotient)));
141 let q = _mm512_or_si512(lo, hi);
142 let product = _mm512_sub_epi32(_mm512_mullo_epi32(value, factor), _mm512_mullo_epi32(q, modulus));
143 let product = _mm512_min_epu32(product, _mm512_sub_epi32(product, modulus));
144 let old = _mm512_loadu_si512(x.as_ptr().add(i).cast());
145 let sum = _mm512_add_epi32(old, product);
146 let sum = _mm512_min_epu32(sum, _mm512_sub_epi32(sum, modulus));
147 _mm512_storeu_si512(x.as_mut_ptr().add(i).cast(), sum);
148 }
149 let a = MInt::new_unchecked(a);
150 for (x, y) in x[end..].iter_mut().zip(&y[end..]) { *x += a * *y; }
151 }
152
153 #[allow(unsafe_op_in_unsafe_fn)]
154 #[target_feature(enable = "avx2")]
155 unsafe fn dot_product_avx2(x: &[MInt<Self>], y: &[MInt<Self>]) -> u32 {
156 use std::arch::x86_64::*;
157 assert_eq!(x.len(), y.len());
158 let modulus = Self::get_mod() as u64;
159 let (products, bound) = if modulus <= 1 << 30 {
160 (8, 2 * modulus)
161 } else if modulus <= 1 << 31 {
162 (2, modulus)
163 } else {
164 (
165 ((u64::MAX - (modulus - 1)) / ((modulus - 1) * (modulus - 1))) as usize,
166 0,
167 )
168 };
169 let bound = _mm256_set1_epi64x((bound << 32) as i64);
171 let len = x.len();
172 let x = x.as_ptr().cast::<u32>();
173 let y = y.as_ptr().cast::<u32>();
174 let mut even = _mm256_setzero_si256();
175 let mut odd = even;
176 let mut offset = 0;
177 let mut result = 0u64;
178 loop {
179 while offset + 8 <= len {
180 let end = (offset + 8 * products).min(len / 8 * 8);
181 while offset < end {
182 let xv = _mm256_loadu_si256(x.add(offset).cast());
183 let yv = _mm256_loadu_si256(y.add(offset).cast());
184 even = _mm256_add_epi64(even, _mm256_mul_epu32(xv, yv));
185 odd = _mm256_add_epi64(
186 odd,
187 _mm256_mul_epu32(_mm256_srli_epi64::<32>(xv), _mm256_srli_epi64::<32>(yv)),
188 );
189 offset += 8;
190 }
191 if modulus > 1 << 31 {
192 break;
193 }
194 even = _mm256_min_epu32(even, _mm256_sub_epi32(even, bound));
195 odd = _mm256_min_epu32(odd, _mm256_sub_epi32(odd, bound));
196 }
197 let mut lanes = [0u64; 8];
198 _mm256_storeu_si256(lanes.as_mut_ptr().cast(), even);
199 _mm256_storeu_si256(lanes.as_mut_ptr().add(4).cast(), odd);
200 let mut low = 0u64;
201 let mut high = 0u64;
202 for lane in lanes {
203 low += lane as u32 as u64;
204 high += lane >> 32;
205 }
206 result = (result + low + (high % modulus) * ((1u64 << 32) % modulus)) % modulus;
207 if offset + 8 > len {
208 break;
209 }
210 even = _mm256_setzero_si256();
211 odd = even;
212 }
213 for first in (offset..len).step_by(products) {
214 for i in first..(first + products).min(len) {
215 result += *x.add(i) as u64 * *y.add(i) as u64;
216 }
217 result %= modulus;
218 }
219 result as u32
220 }
221
222 #[allow(unsafe_op_in_unsafe_fn)]
223 #[target_feature(enable = "avx512f")]
224 unsafe fn dot_product_avx512(x: &[MInt<Self>], y: &[MInt<Self>]) -> u32 {
225 use std::arch::x86_64::*;
226 assert_eq!(x.len(), y.len());
227 let modulus = Self::get_mod() as u64;
228 let (products, bound) = if modulus <= 1 << 30 {
229 (8, 2 * modulus)
230 } else if modulus <= 1 << 31 {
231 (2, modulus)
232 } else {
233 (
234 ((u64::MAX - (modulus - 1)) / ((modulus - 1) * (modulus - 1))) as usize,
235 0,
236 )
237 };
238 let bound = _mm512_set1_epi64((bound << 32) as i64);
240 let len = x.len();
241 let x = x.as_ptr().cast::<u32>();
242 let y = y.as_ptr().cast::<u32>();
243 let mut even = _mm512_setzero_si512();
244 let mut odd = even;
245 let mut offset = 0;
246 let mut result = 0u64;
247 loop {
248 while offset + 16 <= len {
249 let end = (offset + 16 * products).min(len / 16 * 16);
250 while offset < end {
251 let xv = _mm512_loadu_si512(x.add(offset).cast());
252 let yv = _mm512_loadu_si512(y.add(offset).cast());
253 even = _mm512_add_epi64(even, _mm512_mul_epu32(xv, yv));
254 odd = _mm512_add_epi64(
255 odd,
256 _mm512_mul_epu32(_mm512_srli_epi64::<32>(xv), _mm512_srli_epi64::<32>(yv)),
257 );
258 offset += 16;
259 }
260 if modulus > 1 << 31 {
261 break;
262 }
263 even = _mm512_min_epu32(even, _mm512_sub_epi32(even, bound));
264 odd = _mm512_min_epu32(odd, _mm512_sub_epi32(odd, bound));
265 }
266 let mut lanes = [0u64; 16];
267 _mm512_storeu_si512(lanes.as_mut_ptr().cast(), even);
268 _mm512_storeu_si512(lanes.as_mut_ptr().add(8).cast(), odd);
269 let mut low = 0u64;
270 let mut high = 0u64;
271 for lane in lanes {
272 low += lane as u32 as u64;
273 high += lane >> 32;
274 }
275 result = (result + low + (high % modulus) * ((1u64 << 32) % modulus)) % modulus;
276 if offset + 16 > len {
277 break;
278 }
279 even = _mm512_setzero_si512();
280 odd = even;
281 }
282 for first in (offset..len).step_by(products) {
283 for i in first..(first + products).min(len) {
284 result += *x.add(i) as u64 * *y.add(i) as u64;
285 }
286 result %= modulus;
287 }
288 result as u32
289 }
290 }
291 };
292}
293
294impl_basic_mint_dot_product!(u32, u64; Modulo998244353, Modulo1000000007, Modulo1000000009, DynModuloU32);
295impl_basic_mint_dot_product!(u64, u128; DynModuloU64);
296impl MIntDotProduct for Modulo2 {}