Skip to main content

competitive/math/
mint_fft_convolve.rs

1#![allow(unsafe_op_in_unsafe_fn)]
2
3use super::{
4    AssociatedValue, MInt, MIntConvert, advise_huge_pages,
5    fast_fourier_transform::{RotateCache, simd::*},
6};
7use std::arch::x86_64::*;
8
9#[target_feature(enable = "avx2,fma")]
10#[inline]
11unsafe fn round4(value: &[f64; 4]) -> [i64; 4] {
12    let magic = _mm256_set1_pd((3i64 << 51) as f64);
13    let rounded = _mm256_sub_epi64(
14        _mm256_castpd_si256(_mm256_add_pd(_mm256_load_pd(value.as_ptr()), magic)),
15        _mm256_castpd_si256(magic),
16    );
17    let mut result = [0; 4];
18    _mm256_storeu_si256(result.as_mut_ptr().cast(), rounded);
19    result
20}
21
22#[target_feature(enable = "avx2,fma")]
23#[inline]
24unsafe fn reduce_mod4(value: __m256d, modulus: __m256d, inverse: __m256d) -> __m256d {
25    let quotient = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
26        _mm256_mul_pd(value, inverse),
27    );
28    _mm256_fnmadd_pd(quotient, modulus, value)
29}
30
31#[target_feature(enable = "avx2,fma")]
32unsafe fn split_coefficients<M>(
33    values: Vec<MInt<M>>,
34    n: usize,
35    modulus: i64,
36    split: i64,
37) -> (Vec<Complex4>, Vec<Complex4>)
38where
39    M: MIntConvert + MIntConvert<u32>,
40{
41    let mut low = Vec::<Complex4>::with_capacity(n / 4);
42    let mut high = Vec::<Complex4>::with_capacity(n / 4);
43    advise_huge_pages(&mut low);
44    advise_huge_pages(&mut high);
45    let divisor = _mm256_set1_pd(split as f64);
46    let split = _mm256_set1_pd(split as f64);
47    for i in (0..values.len()).step_by(4) {
48        let mut centered = [0i32; 4];
49        for lane in 0..4.min(values.len() - i) {
50            let mut value = u32::from(values[i + lane]) as i64;
51            if value * 2 > modulus {
52                value -= modulus;
53            }
54            centered[lane] = value as i32;
55        }
56        let value = _mm256_cvtepi32_pd(_mm_loadu_si128(centered.as_ptr().cast()));
57        let upper = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
58            _mm256_div_pd(value, divisor),
59        );
60        let lower = _mm256_fnmadd_pd(upper, split, value);
61        if i < n {
62            let mut lo = Complex4::default();
63            let mut hi = Complex4::default();
64            _mm256_store_pd(lo.re.as_mut_ptr(), lower);
65            _mm256_store_pd(hi.re.as_mut_ptr(), upper);
66            low.push(lo);
67            high.push(hi);
68        } else {
69            _mm256_store_pd(low[(i - n) >> 2].im.as_mut_ptr(), lower);
70            _mm256_store_pd(high[(i - n) >> 2].im.as_mut_ptr(), upper);
71        }
72    }
73    low.resize(n / 4, Complex4::default());
74    high.resize(n / 4, Complex4::default());
75    (low, high)
76}
77
78#[target_feature(enable = "avx2,fma")]
79unsafe fn dot_soa(a0: &mut [Complex4], a1: &mut [Complex4], b0: &mut [Complex4], b1: &[Complex4]) {
80    let n = a0.len() * 4;
81    RotateCache::ensure(n / 2);
82    RotateCache::with(|cache| {
83        for i in 0..a0.len() {
84            let (mut cr, mut ci) = load4(&b0[i]);
85            let (mut dr, mut di) = load4(&b1[i]);
86            let mut c0r = _mm256_setzero_pd();
87            let mut c0i = _mm256_setzero_pd();
88            let mut c1r = _mm256_setzero_pd();
89            let mut c1i = _mm256_setzero_pd();
90            let mut c2r = _mm256_setzero_pd();
91            let mut c2i = _mm256_setzero_pd();
92            let w = eval_twiddle(cache, 1, a0.len(), i);
93            let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
94            let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
95            for lane in 0..4 {
96                let ar = _mm256_set1_pd(a0[i].re[lane]);
97                let ai = _mm256_set1_pd(a0[i].im[lane]);
98                let br = _mm256_set1_pd(a1[i].re[lane]);
99                let bi = _mm256_set1_pd(a1[i].im[lane]);
100                multiply_accumulate4(&mut c0r, &mut c0i, ar, ai, cr, ci);
101                multiply_accumulate4(&mut c1r, &mut c1i, ar, ai, dr, di);
102                multiply_accumulate4(&mut c1r, &mut c1i, br, bi, cr, ci);
103                multiply_accumulate4(&mut c2r, &mut c2i, br, bi, dr, di);
104                if lane != 3 {
105                    cr = _mm256_permute4x64_pd::<0x93>(cr);
106                    ci = _mm256_permute4x64_pd::<0x93>(ci);
107                    dr = _mm256_permute4x64_pd::<0x93>(dr);
108                    di = _mm256_permute4x64_pd::<0x93>(di);
109                    (cr, ci) = mul4(cr, ci, wr, wi);
110                    (dr, di) = mul4(dr, di, wr, wi);
111                }
112            }
113            store4(&mut a0[i], c0r, c0i);
114            store4(&mut a1[i], c1r, c1i);
115            store4(&mut b0[i], c2r, c2i);
116        }
117    });
118}
119
120#[target_feature(enable = "avx2,fma")]
121unsafe fn split_u64_coefficients(values: &[u64], n: usize) -> [Vec<Complex4>; 5] {
122    let mut result: [Vec<Complex4>; 5] = std::array::from_fn(|_| {
123        let mut part = Vec::with_capacity(n / 4);
124        advise_huge_pages(&mut part);
125        part
126    });
127    for (i, chunk) in values.chunks(4).enumerate() {
128        let mut parts = [Complex4::default(); 5];
129        for (lane, mut value) in chunk.iter().copied().enumerate() {
130            for part in &mut parts {
131                let digit = ((value << 51) as i64) >> 51;
132                value = (value >> 13).wrapping_add(u64::from(digit < 0));
133                part.re[lane] = digit as f64;
134            }
135        }
136        for (result, part) in result.iter_mut().zip(parts) {
137            if i < n / 4 {
138                result.push(part);
139            } else {
140                result[i - n / 4].im = part.re;
141            }
142        }
143    }
144    for part in &mut result {
145        part.resize(n / 4, Complex4::default());
146    }
147    result
148}
149
150#[target_feature(enable = "avx2,fma")]
151unsafe fn dot_u64_soa(a: &mut [Vec<Complex4>; 5], b: &[Vec<Complex4>; 5]) {
152    let n = a[0].len() * 4;
153    RotateCache::ensure(n / 2);
154    RotateCache::with(|cache| {
155        for block in 0..a[0].len() {
156            let mut br = [_mm256_setzero_pd(); 5];
157            let mut bi = br;
158            let mut rr = br;
159            let mut ri = br;
160            for part in 0..5 {
161                (br[part], bi[part]) = load4(&b[part][block]);
162            }
163            let w = eval_twiddle(cache, 1, a[0].len(), block);
164            let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
165            let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
166            for lane in 0..4 {
167                let ar: [__m256d; 5] =
168                    std::array::from_fn(|part| _mm256_set1_pd(a[part][block].re[lane]));
169                let ai: [__m256d; 5] =
170                    std::array::from_fn(|part| _mm256_set1_pd(a[part][block].im[lane]));
171                for part in 0..5 {
172                    for left in 0..=part {
173                        multiply_accumulate4(
174                            &mut rr[part],
175                            &mut ri[part],
176                            ar[left],
177                            ai[left],
178                            br[part - left],
179                            bi[part - left],
180                        );
181                    }
182                }
183                if lane != 3 {
184                    for part in 0..5 {
185                        br[part] = _mm256_permute4x64_pd::<0x93>(br[part]);
186                        bi[part] = _mm256_permute4x64_pd::<0x93>(bi[part]);
187                        (br[part], bi[part]) = mul4(br[part], bi[part], wr, wi);
188                    }
189                }
190            }
191            for part in 0..5 {
192                store4(&mut a[part][block], rr[part], ri[part]);
193            }
194        }
195    });
196}
197
198#[target_feature(enable = "avx2,fma")]
199pub unsafe fn convolve_u64_avx2(a: Vec<u64>, b: Vec<u64>) -> Vec<u64> {
200    let len = a.len() + b.len() - 1;
201    let n = len.next_power_of_two() / 2;
202    let a_parts = if a.iter().all(|&value| value <= u32::MAX as u64) {
203        3
204    } else {
205        5
206    };
207    let mut fa = split_u64_coefficients(&a, n);
208    drop(a);
209    let b_parts = if b.iter().all(|&value| value <= u32::MAX as u64) {
210        3
211    } else {
212        5
213    };
214    let mut fb = split_u64_coefficients(&b, n);
215    drop(b);
216    for part in 0..3 {
217        fft_soa(&mut fa[part]);
218        fft_soa(&mut fb[part]);
219    }
220    for part in &mut fa[3..a_parts] {
221        fft_soa(part);
222    }
223    for part in &mut fb[3..b_parts] {
224        fft_soa(part);
225    }
226    dot_u64_soa(&mut fa, &fb);
227    drop(fb);
228    for part in &mut fa {
229        ifft_soa(part);
230    }
231    let mut result = vec![0; len];
232    for (block, _) in fa[0].iter().enumerate() {
233        let real: [[i64; 4]; 5] = std::array::from_fn(|part| round4(&fa[part][block].re));
234        let imag: [[i64; 4]; 5] = std::array::from_fn(|part| round4(&fa[part][block].im));
235        for lane in 0..4 {
236            let i = block * 4 + lane;
237            if i < len {
238                result[i] = (real[0][lane] as u64)
239                    .wrapping_add((real[1][lane] as u64) << 13)
240                    .wrapping_add((real[2][lane] as u64) << 26)
241                    .wrapping_add((real[3][lane] as u64) << 39)
242                    .wrapping_add((real[4][lane] as u64) << 52);
243            }
244            if i + n < len {
245                result[i + n] = (imag[0][lane] as u64)
246                    .wrapping_add((imag[1][lane] as u64) << 13)
247                    .wrapping_add((imag[2][lane] as u64) << 26)
248                    .wrapping_add((imag[3][lane] as u64) << 39)
249                    .wrapping_add((imag[4][lane] as u64) << 52);
250            }
251        }
252    }
253    result
254}
255
256#[target_feature(enable = "avx2,fma")]
257pub unsafe fn convolve_mint_avx2<M>(a: Vec<MInt<M>>, b: Vec<MInt<M>>) -> Vec<MInt<M>>
258where
259    M: MIntConvert + MIntConvert<u32>,
260{
261    let len = a.len() + b.len() - 1;
262    let n = len.next_power_of_two() / 2;
263    let modulus = <M as MIntConvert<u32>>::mod_into() as i64;
264    let split = (modulus as f64).sqrt() as i64 + 1;
265    let (mut a0, mut a1) = split_coefficients(a, n, modulus, split);
266    let (mut b0, mut b1) = split_coefficients(b, n, modulus, split);
267    fft_soa(&mut a0);
268    fft_soa(&mut a1);
269    fft_soa(&mut b0);
270    fft_soa(&mut b1);
271    dot_soa(&mut a0, &mut a1, &mut b0, &b1);
272    drop(b1);
273    ifft_soa(&mut a0);
274    ifft_soa(&mut a1);
275    ifft_soa(&mut b0);
276    let split2 = (split * split % modulus) as f64;
277    let split = _mm256_set1_pd(split as f64);
278    let split2 = _mm256_set1_pd(split2);
279    let inverse = _mm256_set1_pd(1.0 / modulus as f64);
280    let modulus = _mm256_set1_pd(modulus as f64);
281    let magic = _mm256_set1_pd((3i64 << 51) as f64);
282    let mut result = vec![MInt::<M>::from(0u32); len];
283    for (block, ((a0, a1), b0)) in a0.iter().zip(&a1).zip(&b0).enumerate() {
284        for (part, (a0, a1, b0)) in [(&a0.re, &a1.re, &b0.re), (&a0.im, &a1.im, &b0.im)]
285            .into_iter()
286            .enumerate()
287        {
288            let a0 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
289                _mm256_load_pd(a0.as_ptr()),
290            );
291            let a1 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
292                _mm256_load_pd(a1.as_ptr()),
293            );
294            let b0 = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(
295                _mm256_load_pd(b0.as_ptr()),
296            );
297            let a0 = reduce_mod4(a0, modulus, inverse);
298            let a1 = reduce_mod4(a1, modulus, inverse);
299            let b0 = reduce_mod4(b0, modulus, inverse);
300            let value = _mm256_fmadd_pd(b0, split2, _mm256_fmadd_pd(a1, split, a0));
301            let value = reduce_mod4(value, modulus, inverse);
302            let value = _mm256_add_pd(
303                value,
304                _mm256_and_pd(
305                    _mm256_cmp_pd::<_CMP_LT_OQ>(value, _mm256_setzero_pd()),
306                    modulus,
307                ),
308            );
309            let value = _mm256_sub_epi64(
310                _mm256_castpd_si256(_mm256_add_pd(value, magic)),
311                _mm256_castpd_si256(magic),
312            );
313            let mut lanes = [0i64; 4];
314            _mm256_storeu_si256(lanes.as_mut_ptr().cast(), value);
315            for (lane, value) in lanes.into_iter().enumerate() {
316                let i = block * 4 + lane + part * n;
317                if i < len {
318                    let value = value as u32;
319                    let modulus = <M as MIntConvert<u32>>::mod_into();
320                    // Expose the reduced range to the conversion's remainder operation.
321                    result[i] = MInt::<M>::from(if value < modulus {
322                        value
323                    } else {
324                        value % modulus
325                    });
326                }
327            }
328        }
329    }
330    result
331}