Skip to main content

fft_soa

Function fft_soa 

Source
pub unsafe fn fft_soa(a: &mut [Complex4])
Examples found in repository?
crates/competitive/src/math/fast_fourier_transform.rs (line 313)
305    pub unsafe fn convolve_f64_avx2(
306        a: impl ExactSizeIterator<Item = f64>,
307        b: impl ExactSizeIterator<Item = f64>,
308        range: std::ops::Range<usize>,
309    ) -> Vec<f64> {
310        let n = (range.end.next_power_of_two() / 2).max(4);
311        let mut fa = pack_f64(a, n);
312        let mut fb = pack_f64(b, n);
313        fft_soa(&mut fa);
314        fft_soa(&mut fb);
315        dot_one_soa(&mut fa, &fb);
316        drop(fb);
317        ifft_soa(&mut fa);
318        range
319            .map(|i| {
320                if i < n {
321                    fa[i >> 2].re[i & 3]
322                } else {
323                    fa[(i - n) >> 2].im[i & 3]
324                }
325            })
326            .collect()
327    }
More examples
Hide additional examples
crates/competitive/src/math/mint_fft_convolve.rs (line 217)
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}