Skip to main content

openvm_stark_backend/dft/
radix_2_bowers_serial.rs

1use std::borrow::BorrowMut;
2
3// Originally copied from p3-dft [src/radix_2_bowers.rs] to turn off rayon
4use 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/// The Bowers G FFT algorithm.
14/// See: "Improved Twiddle Access for Fast Fourier Transforms"
15#[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    /// Compute the inverse DFT of `vec`.
32    fn idft(&self, vec: Vec<F>) -> Vec<F> {
33        self.idft_batch(RowMajorMatrix::new(vec, 1)).values
34    }
35
36    /// Compute the inverse DFT of each column in `mat`.
37    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        // It's cheaper to use div_2exp_u64 as this usually avoids an inversion.
62        // It's also cheaper to work in the PrimeSubfield whenever possible.
63        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        // Rescale coefficients in two ways:
69        // - divide by height (since we're doing an inverse DFT)
70        // - multiply by powers of the coset shift (see default coset LDE impl for an explanation)
71        let weights = Powers {
72            base: shift,
73            current: h_inv,
74        }
75        .take(h);
76        for (row, weight) in weights.enumerate() {
77            // reverse_bits because mat is encoded in bit-reversed order
78            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
89/// Executes the Bowers G network. This is like a DFT, except it assumes the input is in
90/// bit-reversed order.
91fn 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
105/// Executes the Bowers G^T network. This is like an inverse DFT, except we skip rescaling by
106/// 1/height, and the output is bit-reversed.
107fn 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    // It's cheaper to use div_2exp_u64 as this usually avoids an inversion.
149    // It's also cheaper to work in the PrimeSubfield whenever possible.
150    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
182/// Assumes `i < j`.
183///
184/// SAFETY: The caller must ensure `i < j < h`, where `h` is the height of the matrix.
185pub(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}