Skip to main content

competitive/num/mint/
mint_basic_dot_product.rs

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                // This quotient underestimates floor(a*y/m) by at most one.
104                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                // This quotient underestimates floor(a*y/m) by at most one.
135                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                // For moduli up to 2^31, each batch adds less than bound*2^32.
170                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                // For moduli up to 2^31, each batch adds less than bound*2^32.
239                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 {}