pub trait AssociatedValue {
type T: 'static + Clone;
// Required method
unsafe fn __local_key() -> &'static LocalKey<Cell<Self::T>>;
// Provided methods
fn get() -> Self::T { ... }
fn set(x: Self::T) { ... }
fn replace(x: Self::T) -> Self::T { ... }
fn with<F, R>(f: F) -> R
where F: FnOnce(&Self::T) -> R { ... }
fn modify<F, R>(f: F) -> R
where F: FnOnce(&mut Self::T) -> R { ... }
}Expand description
Trait for a modifiable value associated with a type.
Required Associated Types§
Required Methods§
unsafe fn __local_key() -> &'static LocalKey<Cell<Self::T>>
Provided Methods§
fn get() -> Self::T
fn set(x: Self::T)
fn replace(x: Self::T) -> Self::T
Sourcefn with<F, R>(f: F) -> R
fn with<F, R>(f: F) -> R
Examples found in repository?
More examples
crates/competitive/src/math/mint_fft_convolve.rs (lines 82-117)
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}crates/competitive/src/math/fast_fourier_transform.rs (lines 112-180)
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}Sourcefn modify<F, R>(f: F) -> R
fn modify<F, R>(f: F) -> R
Examples found in repository?
crates/competitive/src/math/fast_fourier_transform.rs (lines 11-34)
9 pub fn ensure(n: usize) {
10 assert_eq!(n.count_ones(), 1, "call with power of two but {}", n);
11 Self::modify(|cache| {
12 let mut m = cache.len();
13 assert!(
14 m.count_ones() <= 1,
15 "length might be power of two but {}",
16 m
17 );
18 if m >= n {
19 return;
20 }
21 cache.reserve_exact(n - m);
22 if cache.is_empty() {
23 cache.push(Complex::one());
24 m += 1;
25 }
26 while m < n {
27 let p = Complex::primitive_nth_root_of_unity(-((m * 4) as f64));
28 for i in 0..m {
29 cache.push(cache[i] * p);
30 }
31 m <<= 1;
32 }
33 assert_eq!(cache.len(), n);
34 });
35 }Dyn Compatibility§
This trait is not dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".