pub enum RotateCache {}Implementations§
Source§impl RotateCache
impl RotateCache
Sourcepub fn ensure(n: usize)
pub fn ensure(n: usize)
Examples found in repository?
crates/competitive/src/math/mint_fft_convolve.rs (line 81)
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}More examples
crates/competitive/src/math/fast_fourier_transform.rs (line 111)
109 pub unsafe fn fft_soa(a: &mut [Complex4]) {
110 let n = a.len() * 4;
111 RotateCache::ensure(n / 2);
112 RotateCache::with(|cache| {
113 let parity = n.trailing_zeros() & 1;
114 for leaf in (0..n).step_by(16) {
115 let mut level = (n + leaf).trailing_zeros();
116 level -= u32::from(level & 1 != parity);
117 while level >= 4 {
118 let len = 1usize << level;
119 let q = leaf >> level;
120 let width = len / 16;
121 let start = q * width * 4;
122 let (a, rest) = a[start..start + width * 4].split_at_mut(width);
123 let (b, rest) = rest.split_at_mut(width);
124 let (c, d) = rest.split_at_mut(width);
125 let w1 = eval_twiddle(cache, 4, n >> level, q);
126 let w2 = w1 * w1;
127 let w3 = w1 * w2;
128 let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
129 let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
130 let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
131 for i in 0..width {
132 let (ar, ai) = load4(&a[i]);
133 let (br, bi) = load4(&b[i]);
134 let (cr, ci) = load4(&c[i]);
135 let (dr, di) = load4(&d[i]);
136 let (br, bi) = mul4(br, bi, w1r, w1i);
137 let (cr, ci) = mul4(cr, ci, w2r, w2i);
138 let (dr, di) = mul4(dr, di, w3r, w3i);
139 let acr = _mm256_add_pd(ar, cr);
140 let aci = _mm256_add_pd(ai, ci);
141 let bdr = _mm256_add_pd(br, dr);
142 let bdi = _mm256_add_pd(bi, di);
143 let acd_r = _mm256_sub_pd(ar, cr);
144 let acd_i = _mm256_sub_pd(ai, ci);
145 let bdd_r = _mm256_sub_pd(br, dr);
146 let bdd_i = _mm256_sub_pd(bi, di);
147 store4(&mut a[i], _mm256_add_pd(acr, bdr), _mm256_add_pd(aci, bdi));
148 store4(&mut b[i], _mm256_sub_pd(acr, bdr), _mm256_sub_pd(aci, bdi));
149 store4(
150 &mut c[i],
151 _mm256_sub_pd(acd_r, bdd_i),
152 _mm256_add_pd(acd_i, bdd_r),
153 );
154 store4(
155 &mut d[i],
156 _mm256_add_pd(acd_r, bdd_i),
157 _mm256_sub_pd(acd_i, bdd_r),
158 );
159 }
160 level -= 2;
161 }
162 }
163 if parity != 0 {
164 let blocks = n / 8;
165 for k in 0..blocks {
166 let w = eval_twiddle(cache, 2, blocks, k);
167 let wr = _mm256_set1_pd(w.re);
168 let wi = _mm256_set1_pd(w.im);
169 let (ar, ai) = load4(&a[k * 2]);
170 let (br, bi) = load4(&a[k * 2 + 1]);
171 let (br, bi) = mul4(br, bi, wr, wi);
172 store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
173 store4(
174 &mut a[k * 2 + 1],
175 _mm256_sub_pd(ar, br),
176 _mm256_sub_pd(ai, bi),
177 );
178 }
179 }
180 });
181 }
182
183 #[target_feature(enable = "avx2,fma")]
184 pub unsafe fn ifft_soa(a: &mut [Complex4]) {
185 let n = a.len() * 4;
186 RotateCache::ensure(n / 2);
187 RotateCache::with(|cache| {
188 let parity = n.trailing_zeros() & 1;
189 if parity != 0 {
190 let blocks = n / 8;
191 for k in 0..blocks {
192 let w = eval_twiddle(cache, 2, blocks, k).conjugate();
193 let wr = _mm256_set1_pd(w.re);
194 let wi = _mm256_set1_pd(w.im);
195 let (ar, ai) = load4(&a[k * 2]);
196 let (br, bi) = load4(&a[k * 2 + 1]);
197 store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
198 let (br, bi) = mul4(_mm256_sub_pd(ar, br), _mm256_sub_pd(ai, bi), wr, wi);
199 store4(&mut a[k * 2 + 1], br, bi);
200 }
201 }
202 for leaf in (12..n).step_by(16) {
203 let max_level = (leaf + 3).trailing_ones();
204 let mut level = 4 + parity;
205 while level <= max_level {
206 let len = 1usize << level;
207 let q = leaf >> level;
208 let width = len / 16;
209 let start = q * width * 4;
210 let (a, rest) = a[start..start + width * 4].split_at_mut(width);
211 let (b, rest) = rest.split_at_mut(width);
212 let (c, d) = rest.split_at_mut(width);
213 let w1 = eval_twiddle(cache, 4, n >> level, q).conjugate();
214 let w2 = w1 * w1;
215 let w3 = w1 * w2;
216 let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
217 let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
218 let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
219 for i in 0..width {
220 let (ar, ai) = load4(&a[i]);
221 let (br, bi) = load4(&b[i]);
222 let (cr, ci) = load4(&c[i]);
223 let (dr, di) = load4(&d[i]);
224 let abr = _mm256_add_pd(ar, br);
225 let abi = _mm256_add_pd(ai, bi);
226 let cdr = _mm256_add_pd(cr, dr);
227 let cdi = _mm256_add_pd(ci, di);
228 let abd_r = _mm256_sub_pd(ar, br);
229 let abd_i = _mm256_sub_pd(ai, bi);
230 let cdd_r = _mm256_sub_pd(cr, dr);
231 let cdd_i = _mm256_sub_pd(ci, di);
232 store4(&mut a[i], _mm256_add_pd(abr, cdr), _mm256_add_pd(abi, cdi));
233 let (br, bi) = mul4(
234 _mm256_add_pd(abd_r, cdd_i),
235 _mm256_sub_pd(abd_i, cdd_r),
236 w1r,
237 w1i,
238 );
239 store4(&mut b[i], br, bi);
240 let (cr, ci) =
241 mul4(_mm256_sub_pd(abr, cdr), _mm256_sub_pd(abi, cdi), w2r, w2i);
242 store4(&mut c[i], cr, ci);
243 let (dr, di) = mul4(
244 _mm256_sub_pd(abd_r, cdd_i),
245 _mm256_add_pd(abd_i, cdd_r),
246 w3r,
247 w3i,
248 );
249 store4(&mut d[i], dr, di);
250 }
251 level += 2;
252 }
253 }
254 let scale = _mm256_set1_pd(4.0 / n as f64);
255 for value in a {
256 let (re, im) = load4(value);
257 store4(value, _mm256_mul_pd(re, scale), _mm256_mul_pd(im, scale));
258 }
259 });
260 }
261
262 #[target_feature(enable = "avx2,fma")]
263 unsafe fn dot_one_soa(a: &mut [Complex4], b: &[Complex4]) {
264 let n = a.len() * 4;
265 RotateCache::ensure(n / 2);
266 RotateCache::with(|cache| {
267 for i in 0..a.len() {
268 let (mut br, mut bi) = load4(&b[i]);
269 let mut rr = _mm256_setzero_pd();
270 let mut ri = _mm256_setzero_pd();
271 let w = eval_twiddle(cache, 1, a.len(), i);
272 let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
273 let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
274 for lane in 0..4 {
275 let ar = _mm256_set1_pd(a[i].re[lane]);
276 let ai = _mm256_set1_pd(a[i].im[lane]);
277 multiply_accumulate4(&mut rr, &mut ri, ar, ai, br, bi);
278 if lane != 3 {
279 br = _mm256_permute4x64_pd::<0x93>(br);
280 bi = _mm256_permute4x64_pd::<0x93>(bi);
281 (br, bi) = mul4(br, bi, wr, wi);
282 }
283 }
284 store4(&mut a[i], rr, ri);
285 }
286 });
287 }
288
289 #[inline]
290 fn pack_f64(values: impl Iterator<Item = f64>, n: usize) -> Vec<Complex4> {
291 let mut result = Vec::with_capacity(n / 4);
292 advise_huge_pages(&mut result);
293 result.resize(n / 4, Complex4::default());
294 for (i, value) in values.enumerate() {
295 if i < n {
296 result[i >> 2].re[i & 3] = value;
297 } else {
298 result[(i - n) >> 2].im[i & 3] = value;
299 }
300 }
301 result
302 }
303
304 #[target_feature(enable = "avx2,fma")]
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 }
328 #[target_feature(enable = "avx2")]
329 pub unsafe fn convolve_i64_avx2(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
330 super::convolve_i64_naive(a, b, len)
331 }
332 #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
333 pub unsafe fn convolve_i64_avx512(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
334 super::convolve_i64_naive(a, b, len)
335 }
336}
337
338fn bit_reverse<T>(f: &mut [T]) {
339 let mut ip = vec![0u32];
340 let mut k = f.len();
341 let mut m = 1;
342 while 2 * m < k {
343 k /= 2;
344 for j in 0..m {
345 ip.push(ip[j] + k as u32);
346 }
347 m *= 2;
348 }
349 if m == k {
350 for i in 1..m {
351 for j in 0..i {
352 let ji = j + ip[i] as usize;
353 let ij = i + ip[j] as usize;
354 f.swap(ji, ij);
355 }
356 }
357 } else {
358 for i in 1..m {
359 for j in 0..i {
360 let ji = j + ip[i] as usize;
361 let ij = i + ip[j] as usize;
362 f.swap(ji, ij);
363 f.swap(ji + m, ij + m);
364 }
365 }
366 }
367}
368
369fn real_twiddles(n: usize, inverse: bool, mut f: impl FnMut(usize, Complex<f64>)) {
370 const BLOCK: usize = 256;
371 let sign = if inverse { 1.0 } else { -1.0 };
372 let step = Complex::primitive_nth_root_of_unity(sign * n as f64);
373 for start in (1..n / 4).step_by(BLOCK) {
374 let mut w = Complex::polar(1.0, sign * std::f64::consts::TAU * start as f64 / n as f64);
375 for k in start..(start + BLOCK).min(n / 4) {
376 f(k, w);
377 w *= step;
378 }
379 }
380}
381
382pub fn transform_real(t: impl IntoIterator<Item = f64>, len: usize) -> Vec<Complex<f64>> {
383 let n = len.max(4).next_power_of_two();
384 let mut f = vec![Complex::zero(); n / 2];
385 for (i, t) in t.into_iter().enumerate() {
386 if i & 1 == 0 {
387 f[i / 2].re = t;
388 } else {
389 f[i / 2].im = t;
390 }
391 }
392 fft(&mut f);
393 bit_reverse(&mut f);
394 f[0] = Complex::new(f[0].re + f[0].im, f[0].re - f[0].im);
395 f[n / 4] = f[n / 4].conjugate();
396 real_twiddles(n, false, |k, wk| {
397 let c = wk.conjugate().transpose() + 1.;
398 let d = c * (f[k] - f[n / 2 - k].conjugate()) * 0.5;
399 f[k] -= d;
400 f[n / 2 - k] += d.conjugate();
401 });
402 f
403}
404
405pub fn inverse_transform_real(mut f: Vec<Complex<f64>>, len: usize) -> Vec<f64> {
406 let n = len.max(4).next_power_of_two();
407 assert_eq!(f.len(), n / 2);
408 f[0] = Complex::new((f[0].re + f[0].im) * 0.5, (f[0].re - f[0].im) * 0.5);
409 f[n / 4] = f[n / 4].conjugate();
410 real_twiddles(n, true, |k, wk| {
411 let c = wk.transpose().conjugate() + 1.;
412 let d = c * (f[k] - f[n / 2 - k].conjugate()) * 0.5;
413 f[k] -= d;
414 f[n / 2 - k] += d.conjugate();
415 });
416 bit_reverse(&mut f);
417 ifft(&mut f);
418 let inv = 1. / (n / 2) as f64;
419 (0..len)
420 .map(|i| inv * if i & 1 == 0 { f[i / 2].re } else { f[i / 2].im })
421 .collect()
422}
423
424#[inline(always)]
425fn convolve_i64_naive(a: Vec<i64>, b: Vec<i64>, len: usize) -> Vec<i64> {
426 let (a, b) = if a.len() < b.len() { (b, a) } else { (a, b) };
427 if b.len() == 1 {
428 return a.into_iter().map(|a| a * b[0]).collect();
429 }
430 let mut c = vec![0; len];
431 for (i, a) in a.chunks(1024).enumerate() {
432 for (j, b) in b.iter().enumerate() {
433 let start = i * 1024 + j;
434 for (c, a) in c[start..start + a.len()].iter_mut().zip(a) {
435 *c += *a * *b;
436 }
437 }
438 }
439 c
440}
441
442impl ConvolveSteps for ConvolveRealFft {
443 type T = Vec<i64>;
444 type F = Vec<Complex<f64>>;
445 fn length(t: &Self::T) -> usize {
446 t.len()
447 }
448 fn transform(t: Self::T, len: usize) -> Self::F {
449 transform_real(t.into_iter().map(|t| t as f64), len)
450 }
451 fn inverse_transform(f: Self::F, len: usize) -> Self::T {
452 inverse_transform_real(f, len)
453 .into_iter()
454 .map(|value| value.round() as i64)
455 .collect()
456 }
457 fn convolve(a: Self::T, b: Self::T) -> Self::T {
458 let len = (a.len() + b.len()).saturating_sub(1);
459 // Keep accumulation overflow-free and exact in the FFT's f64 representation.
460 if (a.len().min(b.len()) <= 32 || {
461 let size = len.next_power_of_two();
462 let log = size.ilog2();
463 let limit = crate::avx_helper!(@dispatch simd_backend, SimdBackend;
464 3 * log + 16, 2 * log + 16, 4 * log + 16
465 );
466 2 * a.len() as u128 * b.len() as u128 <= limit as u128 * size as u128
467 }) && a.iter().map(|x| x.unsigned_abs()).max().unwrap_or(0) as u128
468 * b.iter().map(|x| x.unsigned_abs()).max().unwrap_or(0) as u128
469 <= (1u128 << 53) / a.len().min(b.len()).max(1) as u128
470 {
471 return crate::avx_helper!(@dispatch simd_backend, SimdBackend;
472 unsafe {
473 if a.len().max(b.len()) < 32 {
474 simd::convolve_i64_avx2(a, b, len)
475 } else {
476 simd::convolve_i64_avx512(a, b, len)
477 }
478 },
479 unsafe { simd::convolve_i64_avx2(a, b, len) },
480 convolve_i64_naive(a, b, len)
481 );
482 }
483 if !a.is_empty() && !b.is_empty() {
484 crate::avx_helper!(@dispatch_avx2_fma return unsafe {
485 simd::convolve_f64_avx2(
486 a.into_iter().map(|x| x as f64),
487 b.into_iter().map(|x| x as f64),
488 0..len,
489 )
490 .into_iter()
491 .map(|x| x.round() as i64)
492 .collect()
493 }, ());
494 }
495 let mut a = Self::transform(a, len);
496 let b = Self::transform(b, len);
497 Self::multiply(&mut a, &b);
498 Self::inverse_transform(a, len)
499 }
500 fn multiply(f: &mut Self::F, g: &Self::F) {
501 assert_eq!(f.len(), g.len());
502 f[0].re *= g[0].re;
503 f[0].im *= g[0].im;
504 for (f, g) in f.iter_mut().zip(g.iter()).skip(1) {
505 *f *= *g;
506 }
507 }
508}
509
510fn middle_product_f64_scalar(
511 a: impl ExactSizeIterator<Item = f64>,
512 b: impl ExactSizeIterator<Item = f64>,
513) -> Vec<f64> {
514 let a_len = a.len();
515 let b_len = b.len();
516 let len = a_len + b_len - 1;
517 let mut a = transform_real(a, len);
518 let b = transform_real(b, len);
519 ConvolveRealFft::multiply(&mut a, &b);
520 inverse_transform_real(a, len)[b_len - 1..a_len].to_vec()
521}
522
523impl ConvolveRealFft {
524 /// Returns coefficients `b.len() - 1..a.len()` of the convolution of `a` and `b`.
525 /// Panics unless `0 < b.len() <= a.len()`.
526 pub fn middle_product_f64(
527 a: impl ExactSizeIterator<Item = f64>,
528 b: impl ExactSizeIterator<Item = f64>,
529 ) -> Vec<f64> {
530 assert!(0 < b.len() && b.len() <= a.len());
531 crate::avx_helper!(@dispatch_avx2_fma return unsafe {
532 let range = b.len() - 1..a.len();
533 simd::convolve_f64_avx2(a, b, range)
534 }, ());
535 middle_product_f64_scalar(a, b)
536 }
537}
538
539macro_rules! fft_kernel {
540 ($a:expr, $cache:expr, $inverse:expr) => {{
541 let a = $a;
542 let cache = $cache;
543 let n = a.len();
544 if $inverse {
545 let mut v = 1;
546 if n.trailing_zeros() & 1 == 1 {
547 for (a, w) in a.as_chunks_mut::<2>().0.iter_mut().zip(cache) {
548 let y = (a[0] - a[1]) * w.conjugate();
549 a[0] += a[1];
550 a[1] = y;
551 }
552 v = 2;
553 }
554 while v < n {
555 for (q, block) in a.chunks_exact_mut(v * 4).enumerate() {
556 let (a, rest) = block.split_at_mut(v);
557 let (b, rest) = rest.split_at_mut(v);
558 let (c, d) = rest.split_at_mut(v);
559 let w0 = cache[q].conjugate();
560 let w1 = cache[q << 1].conjugate();
561 let w3 = w0 * w1;
562 for (((a, b), c), d) in a.iter_mut().zip(b).zip(c).zip(d) {
563 let ac0 = *a + *b;
564 let ac1 = *c + *d;
565 let bd0 = *a - *b;
566 let bd1 = *c - *d;
567 let bd1 = Complex::new(-bd1.im, bd1.re);
568 *a = ac0 + ac1;
569 *b = (bd0 + bd1) * w1;
570 *c = (ac0 - ac1) * w0;
571 *d = (bd0 - bd1) * w3;
572 }
573 }
574 v <<= 2;
575 }
576 } else {
577 let mut v = n / 2;
578 while v >= 2 {
579 let l = v / 2;
580 for (q, block) in a.chunks_exact_mut(l * 4).enumerate() {
581 let (a, rest) = block.split_at_mut(l);
582 let (b, rest) = rest.split_at_mut(l);
583 let (c, d) = rest.split_at_mut(l);
584 let w0 = cache[q];
585 let w1 = cache[q << 1];
586 let w3 = w0 * w1;
587 for (((a, b), c), d) in a.iter_mut().zip(b).zip(c).zip(d) {
588 let bv = *b * w1;
589 let cv = *c * w0;
590 let dv = *d * w3;
591 let ac0 = *a + cv;
592 let ac1 = *a - cv;
593 let bd0 = bv + dv;
594 let bd1 = bv - dv;
595 let bd1 = Complex::new(bd1.im, -bd1.re);
596 *a = ac0 + bd0;
597 *b = ac0 - bd0;
598 *c = ac1 + bd1;
599 *d = ac1 - bd1;
600 }
601 }
602 v >>= 2;
603 }
604 if v == 1 {
605 for (a, w) in a.as_chunks_mut::<2>().0.iter_mut().zip(cache) {
606 let y = a[1] * *w;
607 a[1] = a[0] - y;
608 a[0] += y;
609 }
610 }
611 }
612 }};
613}
614
615pub fn fft(a: &mut [Complex<f64>]) {
616 fft_dispatch::<false>(a);
617}
618
619pub fn ifft(a: &mut [Complex<f64>]) {
620 fft_dispatch::<true>(a);
621}
622
623fn fft_dispatch<const INVERSE: bool>(a: &mut [Complex<f64>]) {
624 RotateCache::ensure(a.len() / 2);
625 RotateCache::with(|cache| {
626 #[cfg(target_arch = "x86_64")]
627 if a.len() >= 16 {
628 match simd_backend() {
629 SimdBackend::Avx512 => {
630 return unsafe { fft_avx512::<INVERSE>(a, cache) };
631 }
632 SimdBackend::Avx2 => return unsafe { fft_avx2::<INVERSE>(a, cache) },
633 SimdBackend::Scalar => {}
634 }
635 }
636 fft_kernel!(a, cache, INVERSE);
637 });
638}Trait Implementations§
Auto Trait Implementations§
impl Freeze for RotateCache
impl RefUnwindSafe for RotateCache
impl Send for RotateCache
impl Sync for RotateCache
impl Unpin for RotateCache
impl UnsafeUnpin for RotateCache
impl UnwindSafe for RotateCache
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Mutably borrows from an owned value. Read more