openvm_stark_backend/dft/
radix_2_bowers_serial.rs1use std::borrow::BorrowMut;
2
3use p3_dft::{Butterfly, DifButterfly, DitButterfly, TwiddleFreeButterfly, TwoAdicSubgroupDft};
5use p3_field::{Field, PackedValue, Powers, PrimeCharacteristicRing, TwoAdicField};
6use p3_matrix::{
7 dense::{DenseMatrix, DenseStorage, RowMajorMatrix, RowMajorMatrixViewMut},
8 Matrix,
9};
10use p3_util::{log2_strict_usize, reverse_bits, reverse_bits_len, reverse_slice_index_bits};
11use tracing::instrument;
12
13#[derive(Default, Clone)]
16pub struct Radix2BowersSerial;
17
18impl<F: TwoAdicField> TwoAdicSubgroupDft<F> for Radix2BowersSerial {
19 type Evaluations = RowMajorMatrix<F>;
20
21 fn dft(&self, vec: Vec<F>) -> Vec<F> {
22 self.dft_batch(RowMajorMatrix::new_col(vec)).values
23 }
24
25 fn dft_batch(&self, mut mat: RowMajorMatrix<F>) -> RowMajorMatrix<F> {
26 reverse_matrix_index_bits(&mut mat);
27 bowers_g(&mut mat.as_view_mut());
28 mat
29 }
30
31 fn idft(&self, vec: Vec<F>) -> Vec<F> {
33 self.idft_batch(RowMajorMatrix::new(vec, 1)).values
34 }
35
36 fn idft_batch(&self, mut mat: RowMajorMatrix<F>) -> RowMajorMatrix<F> {
38 bowers_g_t(&mut mat.as_view_mut());
39 divide_by_height(&mut mat);
40 reverse_matrix_index_bits(&mut mat);
41 mat
42 }
43
44 fn lde_batch(&self, mut mat: RowMajorMatrix<F>, added_bits: usize) -> RowMajorMatrix<F> {
45 bowers_g_t(&mut mat.as_view_mut());
46 divide_by_height(&mut mat);
47 mat = mat.bit_reversed_zero_pad(added_bits);
48 bowers_g(&mut mat.as_view_mut());
49 mat
50 }
51
52 #[instrument(skip_all, fields(dims = %mat.dimensions(), added_bits))]
53 fn coset_lde_batch(
54 &self,
55 mut mat: RowMajorMatrix<F>,
56 added_bits: usize,
57 shift: F,
58 ) -> RowMajorMatrix<F> {
59 let h = mat.height();
60 let log_h = log2_strict_usize(h);
61 let h_inv_subfield = F::PrimeSubfield::ONE.div_2exp_u64(log_h as u64);
64 let h_inv = F::from_prime_subfield(h_inv_subfield);
65
66 bowers_g_t(&mut mat.as_view_mut());
67
68 let weights = Powers {
72 base: shift,
73 current: h_inv,
74 }
75 .take(h);
76 for (row, weight) in weights.enumerate() {
77 mat.scale_row(reverse_bits(row, h), weight);
79 }
80
81 mat = mat.bit_reversed_zero_pad(added_bits);
82
83 bowers_g(&mut mat.as_view_mut());
84
85 mat
86 }
87}
88
89fn bowers_g<F: TwoAdicField>(mat: &mut RowMajorMatrixViewMut<F>) {
92 let h = mat.height();
93 let log_h = log2_strict_usize(h);
94
95 let root = F::two_adic_generator(log_h);
96 let mut twiddles: Vec<_> = root.powers().take(h / 2).map(DifButterfly).collect();
97 reverse_slice_index_bits(&mut twiddles);
98
99 let log_h = log2_strict_usize(mat.height());
100 for log_half_block_size in 0..log_h {
101 butterfly_layer(mat, 1 << log_half_block_size, &twiddles)
102 }
103}
104
105fn bowers_g_t<F: TwoAdicField>(mat: &mut RowMajorMatrixViewMut<F>) {
108 let h = mat.height();
109 let log_h = log2_strict_usize(h);
110
111 let root_inv = F::two_adic_generator(log_h).inverse();
112 let mut twiddles: Vec<_> = root_inv.powers().take(h / 2).map(DitButterfly).collect();
113 reverse_slice_index_bits(&mut twiddles);
114
115 let log_h = log2_strict_usize(mat.height());
116 for log_half_block_size in (0..log_h).rev() {
117 butterfly_layer(mat, 1 << log_half_block_size, &twiddles)
118 }
119}
120
121fn butterfly_layer<F: Field, B: Butterfly<F>>(
122 mat: &mut RowMajorMatrixViewMut<F>,
123 half_block_size: usize,
124 twiddles: &[B],
125) {
126 mat.row_chunks_exact_mut(2 * half_block_size)
127 .enumerate()
128 .for_each(|(block, mut chunks)| {
129 let (mut hi_chunks, mut lo_chunks) = chunks.split_rows_mut(half_block_size);
130 hi_chunks
131 .rows_mut()
132 .zip(lo_chunks.rows_mut())
133 .for_each(|(hi_chunk, lo_chunk)| {
134 if block == 0 {
135 TwiddleFreeButterfly.apply_to_rows(hi_chunk, lo_chunk)
136 } else {
137 twiddles[block].apply_to_rows(hi_chunk, lo_chunk);
138 }
139 });
140 });
141}
142
143pub fn divide_by_height<F: Field, S: DenseStorage<F> + BorrowMut<[F]>>(
144 mat: &mut DenseMatrix<F, S>,
145) {
146 let h = mat.height();
147 let log_h = log2_strict_usize(h);
148 let h_inv_subfield = F::PrimeSubfield::ONE.div_2exp_u64(log_h as u64);
151 let h_inv = F::from_prime_subfield(h_inv_subfield);
152 scale_slice_in_place(h_inv, mat.values.borrow_mut());
153}
154
155pub fn scale_slice_in_place<F: Field>(s: F, slice: &mut [F]) {
156 let (packed, sfx) = F::Packing::pack_slice_with_suffix_mut(slice);
157 let packed_s: F::Packing = s.into();
158 packed.iter_mut().for_each(|x| *x *= packed_s);
159 sfx.iter_mut().for_each(|x| *x *= s);
160}
161
162#[instrument(level = "debug", skip_all)]
163pub fn reverse_matrix_index_bits<'a, F, S>(mat: &mut DenseMatrix<F, S>)
164where
165 F: Clone + Send + Sync + 'a,
166 S: DenseStorage<F> + BorrowMut<[F]>,
167{
168 let w = mat.width();
169 let h = mat.height();
170 let log_h = log2_strict_usize(h);
171 let values = mat.values.borrow_mut().as_mut_ptr() as usize;
172
173 (0..h).for_each(|i| {
174 let values = values as *mut F;
175 let j = reverse_bits_len(i, log_h);
176 if i < j {
177 unsafe { swap_rows_raw(values, w, i, j) };
178 }
179 });
180}
181
182pub(crate) unsafe fn swap_rows_raw<F>(mat: *mut F, w: usize, i: usize, j: usize) {
186 let row_i = core::slice::from_raw_parts_mut(mat.add(i * w), w);
187 let row_j = core::slice::from_raw_parts_mut(mat.add(j * w), w);
188 row_i.swap_with_slice(row_j);
189}