1use alloc::sync::Arc;
4use alloc::vec::Vec;
5use core::iter;
6
7use itertools::Itertools;
8use p3_field::{Field, TwoAdicField, scale_slice_in_place_single_core};
9use p3_matrix::Matrix;
10use p3_matrix::dense::{RowMajorMatrix, RowMajorMatrixViewMut};
11use p3_matrix::util::reverse_matrix_index_bits;
12use p3_maybe_rayon::prelude::*;
13use p3_util::{as_base_slice, log2_strict_usize, reverse_slice_index_bits};
14use spin::RwLock;
15
16use crate::{
17 Butterfly, DifButterfly, DifButterflyZeros, DitButterfly, TwiddleFreeButterfly,
18 TwoAdicSubgroupDft,
19};
20
21const LAYERS_PER_GROUP: usize = 3;
23
24#[derive(Clone, Debug)]
28struct TwiddlePair<F> {
29 twiddles: Arc<[Vec<F>]>,
30 inv_twiddles: Arc<[Vec<F>]>,
31}
32
33impl<F> Default for TwiddlePair<F> {
34 fn default() -> Self {
35 Self {
36 twiddles: Arc::from(Vec::new()),
37 inv_twiddles: Arc::from(Vec::new()),
38 }
39 }
40}
41
42#[derive(Default, Clone, Debug)]
54pub struct Radix2DFTSmallBatch<F> {
55 cache: Arc<RwLock<TwiddlePair<F>>>,
63}
64
65impl<F: TwoAdicField> Radix2DFTSmallBatch<F> {
66 pub fn new(n: usize) -> Self {
70 let res = Self::default();
71 res.update_twiddles(n);
72 res
73 }
74
75 fn roots_of_unity_table(&self, n: usize) -> Vec<Vec<F>> {
83 let lg_n = log2_strict_usize(n);
84 let generator = F::two_adic_generator(lg_n);
85 let half_n = 1 << (lg_n - 1);
86 let nth_roots = generator.powers().collect_n(half_n);
88
89 (0..lg_n)
90 .map(|i| nth_roots.iter().step_by(1 << i).copied().collect())
91 .collect()
92 }
93
94 fn update_twiddles(&self, fft_len: usize) {
96 let curr_max_fft_len = 1 << self.cache.read().twiddles.len();
102 if fft_len > curr_max_fft_len {
103 let mut new_twiddles = self.roots_of_unity_table(fft_len);
104 let mut new_inv_twiddles: Vec<Vec<F>> = new_twiddles
105 .iter()
106 .map(|ts| {
107 iter::once(F::ONE)
110 .chain(ts[1..].iter().rev().map(|&f| -f))
111 .collect()
112 })
113 .collect();
114
115 new_twiddles.iter_mut().for_each(|ts| {
116 reverse_slice_index_bits(ts);
117 });
118 new_inv_twiddles.iter_mut().for_each(|ts| {
119 reverse_slice_index_bits(ts);
120 });
121
122 let mut cache = self.cache.write();
125 let cur_have = 1usize << cache.twiddles.len();
126 if fft_len > cur_have {
127 cache.twiddles = Arc::from(new_twiddles);
128 cache.inv_twiddles = Arc::from(new_inv_twiddles);
129 }
130 }
131 }
132}
133
134impl<F> TwoAdicSubgroupDft<F> for Radix2DFTSmallBatch<F>
135where
136 F: TwoAdicField,
137{
138 type Evaluations = RowMajorMatrix<F>;
139
140 fn dft_batch(&self, mut mat: RowMajorMatrix<F>) -> Self::Evaluations {
141 let h = mat.height();
142 let w = mat.width();
143 let log_h = log2_strict_usize(h);
144
145 self.update_twiddles(h);
146 let g = self.cache.read().twiddles.clone(); let root_table = &g[g.len() - log_h..];
148
149 let num_par_rows = estimate_num_rows_in_l1::<F>(h, w);
154 let log_num_par_rows = log2_strict_usize(num_par_rows);
155 let chunk_size = num_par_rows * w;
156
157 let multi_layer_dit = MultiLayerDitButterfly {};
161
162 for (dit_0, dit_1, dit_2) in root_table[log_num_par_rows..]
165 .iter()
166 .rev()
167 .map(|slice| unsafe { as_base_slice::<DitButterfly<F>, F>(slice) }) .tuples()
169 {
170 dft_layer_par_triple(&mut mat.as_view_mut(), dit_0, dit_1, dit_2, multi_layer_dit);
171 }
172
173 let corr = (log_h - log_num_par_rows) % LAYERS_PER_GROUP;
176 let extra_layers: Vec<&[DitButterfly<F>]> = root_table
177 [log_num_par_rows..log_num_par_rows + corr]
178 .iter()
179 .map(|slice| unsafe { as_base_slice::<DitButterfly<F>, F>(slice) }) .collect();
181 dft_layer_par_extra_layers(&mut mat.as_view_mut(), &extra_layers, multi_layer_dit);
182
183 par_remaining_layers(&mut mat.values, chunk_size, &root_table[..log_num_par_rows]);
187
188 reverse_matrix_index_bits(&mut mat);
190 mat
191 }
192
193 fn idft_batch(&self, mut mat: RowMajorMatrix<F>) -> RowMajorMatrix<F> {
194 let h = mat.height();
195 let w = mat.width();
196 let log_h = log2_strict_usize(h);
197
198 self.update_twiddles(h);
199 let g = self.cache.read().inv_twiddles.clone(); let start = g
201 .len()
202 .checked_sub(log_h)
203 .expect("log_h exceeds inv_twiddles length");
204 let root_table = &g[start..];
205
206 let num_par_rows = estimate_num_rows_in_l1::<F>(h, w);
212 let log_num_par_rows = log2_strict_usize(num_par_rows);
213 let chunk_size = num_par_rows * w;
214
215 reverse_matrix_index_bits(&mut mat);
217
218 par_initial_layers(
224 &mut mat.values,
225 chunk_size,
226 &root_table[..log_num_par_rows],
227 log_h,
228 );
229
230 let multi_layer_dif = MultiLayerDifButterfly {};
234
235 let corr = (log_h - log_num_par_rows) % LAYERS_PER_GROUP;
238 let extra_layers: Vec<&[DifButterfly<F>]> = root_table
239 [log_num_par_rows..log_num_par_rows + corr]
240 .iter()
241 .map(|slice| unsafe { as_base_slice::<DifButterfly<F>, F>(slice) }) .collect();
243 dft_layer_par_extra_layers(&mut mat.as_view_mut(), &extra_layers, multi_layer_dif);
244
245 for (dif_0, dif_1, dif_2) in root_table[(log_num_par_rows + corr)..]
248 .iter()
249 .map(|slice| unsafe { as_base_slice::<DifButterfly<F>, F>(slice) }) .tuples()
251 {
252 dft_layer_par_triple(&mut mat.as_view_mut(), dif_2, dif_1, dif_0, multi_layer_dif);
253 }
254
255 mat
256 }
257
258 fn coset_lde_batch(
259 &self,
260 mut mat: RowMajorMatrix<F>,
261 added_bits: usize,
262 shift: F,
263 ) -> Self::Evaluations {
264 let h = mat.height();
265 let w = mat.width();
266 let log_h = log2_strict_usize(h);
267
268 self.update_twiddles(h << added_bits);
269 let cached = self.cache.read().clone();
270 let g = &cached.twiddles;
271 let start = g
272 .len()
273 .checked_sub(log_h + added_bits)
274 .expect("log_h exceeds twiddles length");
275 let root_table = &g[start..];
276 let ig = &cached.inv_twiddles;
277 let start = ig
278 .len()
279 .checked_sub(log_h)
280 .expect("log_h exceeds inv_twiddles length");
281 let inv_root_table = &ig[start..];
282 let output_height = h << added_bits;
283
284 let output_values = F::zero_vec(output_height * w);
286 let mut out = RowMajorMatrix::new(output_values, w);
287
288 let num_par_rows = estimate_num_rows_in_l1::<F>(h, w);
305 let num_inner_dit_layers = log2_strict_usize(num_par_rows);
306 let num_inner_dif_layers = num_inner_dit_layers + added_bits;
307
308 let multi_layer_dit = MultiLayerDitButterfly {};
311 for (dit_0, dit_1, dit_2) in inv_root_table[num_inner_dit_layers..]
312 .iter()
313 .rev()
314 .map(|slice| unsafe { as_base_slice::<DitButterfly<F>, F>(slice) }) .tuples()
316 {
317 dft_layer_par_triple(&mut mat.as_view_mut(), dit_0, dit_1, dit_2, multi_layer_dit);
318 }
319
320 let corr = (log_h - num_inner_dit_layers) % LAYERS_PER_GROUP;
323 let extra_layers: Vec<&[DitButterfly<F>]> = inv_root_table
324 [num_inner_dit_layers..num_inner_dit_layers + corr]
325 .iter()
326 .map(|slice| unsafe { as_base_slice::<DitButterfly<F>, F>(slice) }) .collect();
328 dft_layer_par_extra_layers(&mut mat.as_view_mut(), &extra_layers, multi_layer_dit);
329
330 par_middle_layers(
334 &mut mat.as_view_mut(),
335 &mut out.as_view_mut(),
336 num_par_rows,
337 &root_table[..(num_inner_dif_layers)],
338 &inv_root_table[..num_inner_dit_layers],
339 added_bits,
340 shift,
341 );
342
343 let multi_layer_dif = MultiLayerDifButterfly {};
345
346 let extra_layers: Vec<&[DifButterfly<F>]> = root_table
349 [num_inner_dif_layers..num_inner_dif_layers + corr]
350 .iter()
351 .map(|slice| unsafe { as_base_slice::<DifButterfly<F>, F>(slice) }) .collect();
353 dft_layer_par_extra_layers(&mut out.as_view_mut(), &extra_layers, multi_layer_dif);
354
355 for (dif_0, dif_1, dif_2) in root_table[(num_inner_dif_layers + corr)..]
358 .iter()
359 .map(|slice| unsafe { as_base_slice::<DifButterfly<F>, F>(slice) }) .tuples()
361 {
362 dft_layer_par_triple(&mut out.as_view_mut(), dif_2, dif_1, dif_0, multi_layer_dif);
363 }
364
365 out
366 }
367}
368
369#[inline]
382fn dft_layer_par<F: Field, B: Butterfly<F>>(
383 mat: &mut RowMajorMatrixViewMut<'_, F>,
384 twiddles: &[B],
385) {
386 debug_assert!(
387 mat.height().is_multiple_of(twiddles.len()),
388 "Matrix height must be divisible by the number of twiddles"
389 );
390 let size = mat.values.len();
391 let num_blocks = twiddles.len();
392
393 let outer_block_size = size / num_blocks;
394 let half_outer_block_size = outer_block_size / 2;
395
396 mat.values
397 .par_chunks_exact_mut(outer_block_size)
398 .enumerate()
399 .for_each(|(ind, block)| {
400 let (hi_chunk, lo_chunk) = block.split_at_mut(half_outer_block_size);
402
403 let num_threads = current_num_threads();
405 let inner_block_size = size / (2 * num_blocks).max(num_threads);
406
407 hi_chunk
408 .par_chunks_mut(inner_block_size)
409 .zip(lo_chunk.par_chunks_mut(inner_block_size))
410 .for_each(|(hi_chunk, lo_chunk)| {
411 if ind == 0 {
412 TwiddleFreeButterfly.apply_to_rows(hi_chunk, lo_chunk);
414 } else {
415 twiddles[ind].apply_to_rows(hi_chunk, lo_chunk);
417 }
418 });
419 });
420}
421
422#[inline]
427fn par_remaining_layers<F: Field>(mat: &mut [F], chunk_size: usize, root_table: &[Vec<F>]) {
428 mat.par_chunks_exact_mut(chunk_size)
429 .enumerate()
430 .for_each(|(index, chunk)| {
431 remaining_layers(chunk, root_table, index);
432 });
433}
434
435fn remaining_layers<F: Field>(chunk: &mut [F], root_table: &[Vec<F>], index: usize) {
437 for (layer, twiddles) in root_table.iter().rev().enumerate() {
438 let num_twiddles_per_block = 1 << layer;
439 let start = index * num_twiddles_per_block;
440 let twiddle_range = start..(start + num_twiddles_per_block);
441 let dit_twiddles: &[DitButterfly<F>] = unsafe { as_base_slice(&twiddles[twiddle_range]) };
443 dft_layer(chunk, dit_twiddles);
444 }
445}
446
447#[inline]
455fn par_initial_layers<F: Field>(
456 mat: &mut [F],
457 chunk_size: usize,
458 root_table: &[Vec<F>],
459 log_height: usize,
460) {
461 let inv_height = F::ONE.div_2exp_u64(log_height as u64);
462 mat.par_chunks_exact_mut(chunk_size)
463 .enumerate()
464 .for_each(|(index, chunk)| {
465 scale_slice_in_place_single_core(chunk, inv_height);
467 initial_layers(chunk, root_table, index);
468 });
469}
470
471#[inline]
473fn initial_layers<F: Field>(chunk: &mut [F], root_table: &[Vec<F>], index: usize) {
474 let num_rounds = root_table.len();
475
476 for (layer, twiddles) in root_table.iter().enumerate() {
477 let num_twiddles_per_block = 1 << (num_rounds - layer - 1);
478 let start = index * num_twiddles_per_block;
479 let twiddle_range = start..(start + num_twiddles_per_block);
480 let dif_twiddles: &[DifButterfly<F>] = unsafe { as_base_slice(&twiddles[twiddle_range]) };
482 dft_layer(chunk, dif_twiddles);
483 }
484}
485
486fn par_middle_layers<F: Field>(
492 in_mat: &mut RowMajorMatrixViewMut<'_, F>,
493 out_mat: &mut RowMajorMatrixViewMut<'_, F>,
494 num_par_rows: usize,
495 root_table: &[Vec<F>],
496 inv_root_table: &[Vec<F>],
497 added_bits: usize,
498 shift: F,
499) {
500 debug_assert_eq!(in_mat.width(), out_mat.width());
501 debug_assert_eq!(in_mat.height() << added_bits, out_mat.height());
502
503 let width = in_mat.width();
504 let height = in_mat.height();
505 let num_rounds = root_table.len();
506 let in_chunk_size = num_par_rows * width;
507 let out_chunk_size = in_chunk_size << added_bits;
508
509 let log_height = log2_strict_usize(height);
510 let inv_height = F::ONE.div_2exp_u64(log_height as u64);
511
512 let mut scaling = shift.shifted_powers(inv_height).collect_n(height);
513 reverse_slice_index_bits(&mut scaling);
514
515 in_mat
516 .values
517 .par_chunks_exact_mut(in_chunk_size)
518 .zip(out_mat.values.par_chunks_exact_mut(out_chunk_size))
519 .zip(scaling.par_chunks_exact_mut(num_par_rows))
520 .enumerate()
521 .for_each(|(index, ((in_chunk, out_chunk), scaling))| {
522 remaining_layers(in_chunk, inv_root_table, index);
523
524 in_chunk
526 .chunks_exact(width)
527 .zip(scaling)
528 .zip(out_chunk.chunks_exact_mut(width << added_bits))
529 .for_each(|((in_row, scale), out_row)| {
530 out_row
531 .iter_mut()
532 .zip(in_row.iter())
533 .for_each(|(out_val, in_val)| {
534 *out_val = *in_val * *scale;
535 });
536 });
537
538 for (layer, twiddles) in root_table[..added_bits].iter().enumerate() {
541 let num_twiddles_per_block = 1 << (num_rounds - layer - 1);
542 let start = index * num_twiddles_per_block;
543 let twiddle_range = start..(start + num_twiddles_per_block);
544
545 let dif_twiddles_zeros: &[DifButterflyZeros<F>] =
547 unsafe { as_base_slice(&twiddles[twiddle_range]) };
548 dft_layer_zeros(out_chunk, dif_twiddles_zeros, added_bits - layer - 1);
549 }
550
551 initial_layers(out_chunk, &root_table[added_bits..], index);
552 });
553}
554
555#[inline]
564fn dft_layer<F: Field, B: Butterfly<F>>(vec: &mut [F], twiddles: &[B]) {
565 debug_assert_eq!(
566 vec.len() % twiddles.len(),
567 0,
568 "Vector length must be divisible by the number of twiddles"
569 );
570 let size = vec.len();
571 let num_blocks = twiddles.len();
572
573 let block_size = size / num_blocks;
574 let half_block_size = block_size / 2;
575
576 vec.chunks_exact_mut(block_size)
577 .zip(twiddles)
578 .for_each(|(block, &twiddle)| {
579 let (hi_chunk, lo_chunk) = block.split_at_mut(half_block_size);
581
582 twiddle.apply_to_rows(hi_chunk, lo_chunk);
584 });
585}
586
587#[inline]
599fn dft_layer_par_double<F: Field, B: Butterfly<F>, M: MultiLayerButterfly<F, B>>(
600 mat: &mut RowMajorMatrixViewMut<'_, F>,
601 twiddles_small: &[B],
602 twiddles_large: &[B],
603 multi_butterfly: M,
604) {
605 debug_assert!(
606 mat.height().is_multiple_of(twiddles_small.len()),
607 "Matrix height must be divisible by the number of twiddles"
608 );
609 let size = mat.values.len();
610 let num_blocks = twiddles_small.len();
611
612 let outer_block_size = size / num_blocks;
613 let quarter_outer_block_size = outer_block_size / 4;
614
615 let inner_chunk_size =
618 (workload_size::<F>().next_power_of_two() / 4).min(quarter_outer_block_size);
619
620 mat.values
621 .par_chunks_exact_mut(outer_block_size)
622 .enumerate()
623 .for_each(|(ind, block)| {
624 let chunk_par_iters_0 = block
627 .chunks_exact_mut(quarter_outer_block_size)
628 .map(|chunk| chunk.par_chunks_mut(inner_chunk_size))
629 .collect::<Vec<_>>();
630 let chunk_par_iters_1 = zip_par_iter_vec(chunk_par_iters_0);
631 chunk_par_iters_1.into_iter().tuples().for_each(|(hi, lo)| {
632 hi.zip(lo).for_each(|chunks| {
633 multi_butterfly.apply_2_layers(chunks, ind, twiddles_small, twiddles_large);
634 });
635 });
636 });
637}
638
639#[inline]
652fn dft_layer_par_triple<F: Field, B: Butterfly<F>, M: MultiLayerButterfly<F, B>>(
653 mat: &mut RowMajorMatrixViewMut<'_, F>,
654 twiddles_small: &[B],
655 twiddles_med: &[B],
656 twiddles_large: &[B],
657 multi_butterfly: M,
658) {
659 debug_assert!(
660 mat.height().is_multiple_of(twiddles_small.len()),
661 "Matrix height must be divisible by the number of twiddles"
662 );
663 let size = mat.values.len();
664 let num_blocks = twiddles_small.len();
665
666 let outer_block_size = size / num_blocks;
667 let eighth_outer_block_size = outer_block_size / 8;
668
669 let inner_chunk_size =
672 (workload_size::<F>().next_power_of_two() / 8).min(eighth_outer_block_size);
673
674 mat.values
675 .par_chunks_exact_mut(outer_block_size)
676 .enumerate()
677 .for_each(|(ind, block)| {
678 let chunk_par_iters_0 = block
681 .chunks_exact_mut(eighth_outer_block_size)
682 .map(|chunk| chunk.par_chunks_mut(inner_chunk_size))
683 .collect::<Vec<_>>();
684 let chunk_par_iters_1 = zip_par_iter_vec(chunk_par_iters_0);
685 let chunk_par_iters_2 = zip_par_iter_vec(chunk_par_iters_1);
686 chunk_par_iters_2.into_iter().tuples().for_each(|(hi, lo)| {
687 hi.zip(lo).for_each(|chunks| {
688 multi_butterfly.apply_3_layers(
689 chunks,
690 ind,
691 twiddles_small,
692 twiddles_med,
693 twiddles_large,
694 );
695 });
696 });
697 });
698}
699
700fn dft_layer_par_extra_layers<F: Field, B: Butterfly<F>, M: MultiLayerButterfly<F, B>>(
705 mat: &mut RowMajorMatrixViewMut<'_, F>,
706 root_table: &[&[B]],
707 multi_layer: M,
708) {
709 match root_table.len() {
710 1 => {
711 dft_layer_par(&mut mat.as_view_mut(), root_table[0]);
712 }
713 2 => {
714 dft_layer_par_double(
715 &mut mat.as_view_mut(),
716 root_table[1],
717 root_table[0],
718 multi_layer,
719 );
720 }
721 0 => {}
722 _ => unreachable!("The number of layers must be 0, 1 or 2"),
723 }
724}
725
726#[inline]
749fn dft_layer_zeros<F: Field, B: Butterfly<F>>(vec: &mut [F], twiddles: &[B], skip: usize) {
750 debug_assert_eq!(
751 vec.len() % twiddles.len(),
752 0,
753 "Vector length must be divisible by the number of twiddles"
754 );
755 let size = vec.len();
756 let num_blocks = twiddles.len();
757
758 let block_size = size / num_blocks;
759 let half_block_size = block_size / 2;
760
761 vec.chunks_exact_mut(block_size)
762 .zip(twiddles)
763 .step_by(1 << skip) .for_each(|(block, &twiddle)| {
765 let (hi_chunk, lo_chunk) = block.split_at_mut(half_block_size);
767
768 twiddle.apply_to_rows(hi_chunk, lo_chunk);
770 });
771}
772
773type DoubleLayerBlockDecomposition<'a, F> =
775 ((&'a mut [F], &'a mut [F]), (&'a mut [F], &'a mut [F]));
776
777#[inline]
779fn fft_double_layer_single_twiddle<F: Field, Fly: Butterfly<F>>(
780 block: &mut DoubleLayerBlockDecomposition<'_, F>,
781 butterfly: Fly,
782) {
783 butterfly.apply_to_rows(block.0.0, block.1.0);
784 butterfly.apply_to_rows(block.0.1, block.1.1);
785}
786
787#[inline]
792fn fft_double_layer_double_twiddle<F: Field, Fly0: Butterfly<F>, Fly1: Butterfly<F>>(
793 block: &mut DoubleLayerBlockDecomposition<'_, F>,
794 fly0: Fly0,
795 fly1: Fly1,
796) {
797 fly0.apply_to_rows(block.0.0, block.0.1);
798 fly1.apply_to_rows(block.1.0, block.1.1);
799}
800
801type TripleLayerBlockDecomposition<'a, F> = (
803 ((&'a mut [F], &'a mut [F]), (&'a mut [F], &'a mut [F])),
804 ((&'a mut [F], &'a mut [F]), (&'a mut [F], &'a mut [F])),
805);
806
807#[inline]
809fn fft_triple_layer_single_twiddle<F: Field, Fly: Butterfly<F>>(
810 block: &mut TripleLayerBlockDecomposition<'_, F>,
811 butterfly: Fly,
812) {
813 butterfly.apply_to_rows(block.0.0.0, block.1.0.0);
814 butterfly.apply_to_rows(block.0.0.1, block.1.0.1);
815 butterfly.apply_to_rows(block.0.1.0, block.1.1.0);
816 butterfly.apply_to_rows(block.0.1.1, block.1.1.1);
817}
818
819#[inline]
824fn fft_triple_layer_double_twiddle<F: Field, Fly0: Butterfly<F>, Fly1: Butterfly<F>>(
825 block: &mut TripleLayerBlockDecomposition<'_, F>,
826 fly0: Fly0,
827 fly1: Fly1,
828) {
829 fly0.apply_to_rows(block.0.0.0, block.0.1.0);
830 fly0.apply_to_rows(block.0.0.1, block.0.1.1);
831 fly1.apply_to_rows(block.1.0.0, block.1.1.0);
832 fly1.apply_to_rows(block.1.0.1, block.1.1.1);
833}
834
835#[inline]
840fn fft_triple_layer_quad_twiddle<F: Field, Fly0: Butterfly<F>, Flies: Butterfly<F>>(
841 block: &mut TripleLayerBlockDecomposition<'_, F>,
842 fly0: Fly0,
843 butterflies: &[Flies],
844) {
845 debug_assert!(butterflies.len() == 3);
846 fly0.apply_to_rows(block.0.0.0, block.0.0.1);
847 butterflies[0].apply_to_rows(block.0.1.0, block.0.1.1);
848 butterflies[1].apply_to_rows(block.1.0.0, block.1.0.1);
849 butterflies[2].apply_to_rows(block.1.1.0, block.1.1.1);
850}
851
852#[must_use]
857const fn workload_size<T: Sized>() -> usize {
858 const L1_CACHE_SIZE: usize = 1 << 15; L1_CACHE_SIZE / size_of::<T>()
860}
861
862#[must_use]
868fn estimate_num_rows_in_l1<T: Sized>(height: usize, width: usize) -> usize {
869 (workload_size::<T>() / width)
870 .next_power_of_two()
871 .min(height) }
873
874#[inline]
881fn zip_par_iter_vec<I: IndexedParallelIterator>(
882 in_vec: Vec<I>,
883) -> Vec<impl IndexedParallelIterator<Item = (I::Item, I::Item)>> {
884 in_vec
885 .into_iter()
886 .tuples()
887 .map(|(hi, lo)| hi.zip(lo))
888 .collect::<Vec<_>>()
889}
890
891trait MultiLayerButterfly<F: Field, B: Butterfly<F>>: Copy + Send + Sync {
892 fn apply_2_layers(
893 &self,
894 chunk_decomposition: DoubleLayerBlockDecomposition<'_, F>,
895 ind: usize,
896 twiddles_small: &[B],
897 twiddles_large: &[B],
898 );
899
900 fn apply_3_layers(
901 &self,
902 chunk_decomposition: TripleLayerBlockDecomposition<'_, F>,
903 ind: usize,
904 twiddles_small: &[B],
905 twiddles_med: &[B],
906 twiddles_large: &[B],
907 );
908}
909
910#[derive(Debug, Clone, Copy)]
911struct MultiLayerDitButterfly;
912
913impl<F: Field> MultiLayerButterfly<F, DitButterfly<F>> for MultiLayerDitButterfly {
914 #[inline]
915 fn apply_2_layers(
916 &self,
917 mut blk_decomp: DoubleLayerBlockDecomposition<'_, F>,
918 ind: usize,
919 twiddles_small: &[DitButterfly<F>],
920 twiddles_large: &[DitButterfly<F>],
921 ) {
922 if ind == 0 {
923 fft_double_layer_single_twiddle(&mut blk_decomp, TwiddleFreeButterfly);
924 fft_double_layer_double_twiddle(
925 &mut blk_decomp,
926 TwiddleFreeButterfly,
927 twiddles_large[1],
928 );
929 } else {
930 fft_double_layer_single_twiddle(&mut blk_decomp, twiddles_small[ind]);
931 fft_double_layer_double_twiddle(
932 &mut blk_decomp,
933 twiddles_large[2 * ind],
934 twiddles_large[2 * ind + 1],
935 );
936 }
937 }
938
939 #[inline]
940 fn apply_3_layers(
941 &self,
942 mut blk_decomp: TripleLayerBlockDecomposition<'_, F>,
943 ind: usize,
944 twiddles_small: &[DitButterfly<F>],
945 twiddles_med: &[DitButterfly<F>],
946 twiddles_large: &[DitButterfly<F>],
947 ) {
948 if ind == 0 {
949 fft_triple_layer_single_twiddle(&mut blk_decomp, TwiddleFreeButterfly);
950 fft_triple_layer_double_twiddle(&mut blk_decomp, TwiddleFreeButterfly, twiddles_med[1]);
951 fft_triple_layer_quad_twiddle(
952 &mut blk_decomp,
953 TwiddleFreeButterfly,
954 &twiddles_large[1..4],
955 );
956 } else {
957 fft_triple_layer_single_twiddle(&mut blk_decomp, twiddles_small[ind]);
958 fft_triple_layer_double_twiddle(
959 &mut blk_decomp,
960 twiddles_med[2 * ind],
961 twiddles_med[2 * ind + 1],
962 );
963 fft_triple_layer_quad_twiddle(
964 &mut blk_decomp,
965 twiddles_large[4 * ind],
966 &twiddles_large[4 * ind + 1..4 * (ind + 1)],
967 );
968 }
969 }
970}
971
972#[derive(Debug, Clone, Copy)]
973struct MultiLayerDifButterfly;
974
975impl<F: Field> MultiLayerButterfly<F, DifButterfly<F>> for MultiLayerDifButterfly {
976 #[inline]
977 fn apply_2_layers(
978 &self,
979 mut blk_decomp: DoubleLayerBlockDecomposition<'_, F>,
980 ind: usize,
981 twiddles_small: &[DifButterfly<F>],
982 twiddles_large: &[DifButterfly<F>],
983 ) {
984 if ind == 0 {
985 fft_double_layer_double_twiddle(
986 &mut blk_decomp,
987 TwiddleFreeButterfly,
988 twiddles_large[1],
989 );
990 fft_double_layer_single_twiddle(&mut blk_decomp, TwiddleFreeButterfly);
991 } else {
992 fft_double_layer_double_twiddle(
993 &mut blk_decomp,
994 twiddles_large[2 * ind],
995 twiddles_large[2 * ind + 1],
996 );
997 fft_double_layer_single_twiddle(&mut blk_decomp, twiddles_small[ind]);
998 }
999 }
1000
1001 #[inline]
1002 fn apply_3_layers(
1003 &self,
1004 mut blk_decomp: TripleLayerBlockDecomposition<'_, F>,
1005 ind: usize,
1006 twiddles_small: &[DifButterfly<F>],
1007 twiddles_med: &[DifButterfly<F>],
1008 twiddles_large: &[DifButterfly<F>],
1009 ) {
1010 if ind == 0 {
1011 fft_triple_layer_quad_twiddle(
1012 &mut blk_decomp,
1013 TwiddleFreeButterfly,
1014 &twiddles_large[1..4],
1015 );
1016 fft_triple_layer_double_twiddle(&mut blk_decomp, TwiddleFreeButterfly, twiddles_med[1]);
1017 fft_triple_layer_single_twiddle(&mut blk_decomp, TwiddleFreeButterfly);
1018 } else {
1019 fft_triple_layer_quad_twiddle(
1020 &mut blk_decomp,
1021 twiddles_large[4 * ind],
1022 &twiddles_large[4 * ind + 1..4 * (ind + 1)],
1023 );
1024 fft_triple_layer_double_twiddle(
1025 &mut blk_decomp,
1026 twiddles_med[2 * ind],
1027 twiddles_med[2 * ind + 1],
1028 );
1029 fft_triple_layer_single_twiddle(&mut blk_decomp, twiddles_small[ind]);
1030 }
1031 }
1032}