use alloc::vec::Vec;
use p3_field::{BasedVectorSpace, TwoAdicField};
use p3_matrix::Matrix;
use p3_matrix::bitrev::BitReversibleMatrix;
use p3_matrix::dense::{RowMajorMatrix, RowMajorMatrixViewMut};
use p3_matrix::util::swap_rows;
use crate::util::{coset_shift_cols, divide_by_height};
pub trait TwoAdicSubgroupDft<F: TwoAdicField>: Clone + Default {
type Evaluations: BitReversibleMatrix<F> + 'static;
fn dft(&self, vec: Vec<F>) -> Vec<F> {
self.dft_batch(RowMajorMatrix::new_col(vec))
.to_row_major_matrix()
.values
}
fn dft_batch(&self, mat: RowMajorMatrix<F>) -> Self::Evaluations;
fn coset_dft(&self, vec: Vec<F>, shift: F) -> Vec<F> {
self.coset_dft_batch(RowMajorMatrix::new_col(vec), shift)
.to_row_major_matrix()
.values
}
fn coset_dft_batch(&self, mut mat: RowMajorMatrix<F>, shift: F) -> Self::Evaluations {
coset_shift_cols(&mut mat, shift);
self.dft_batch(mat)
}
fn idft(&self, vec: Vec<F>) -> Vec<F> {
self.idft_batch(RowMajorMatrix::new_col(vec)).values
}
fn idft_batch(&self, mat: RowMajorMatrix<F>) -> RowMajorMatrix<F> {
let mut dft = self.dft_batch(mat).to_row_major_matrix();
let h = dft.height();
divide_by_height(&mut dft);
for row in 1..h / 2 {
swap_rows(&mut dft, row, h - row);
}
dft
}
fn coset_idft(&self, vec: Vec<F>, shift: F) -> Vec<F> {
self.coset_idft_batch(RowMajorMatrix::new_col(vec), shift)
.values
}
fn coset_idft_batch(&self, mut mat: RowMajorMatrix<F>, shift: F) -> RowMajorMatrix<F> {
mat = self.idft_batch(mat);
coset_shift_cols(&mut mat, shift.inverse());
mat
}
fn lde(&self, vec: Vec<F>, added_bits: usize) -> Vec<F> {
self.lde_batch(RowMajorMatrix::new_col(vec), added_bits)
.to_row_major_matrix()
.values
}
fn lde_batch(&self, mat: RowMajorMatrix<F>, added_bits: usize) -> Self::Evaluations {
self.coset_lde_batch(mat, added_bits, F::ONE)
}
fn coset_lde(&self, vec: Vec<F>, added_bits: usize, shift: F) -> Vec<F> {
self.coset_lde_batch(RowMajorMatrix::new_col(vec), added_bits, shift)
.to_row_major_matrix()
.values
}
fn coset_lde_batch(
&self,
mat: RowMajorMatrix<F>,
added_bits: usize,
shift: F,
) -> Self::Evaluations {
self.coset_lde_batch_with_transform(mat, added_bits, shift, |_, _| {})
}
fn coset_lde_batch_with_transform<T>(
&self,
mat: RowMajorMatrix<F>,
added_bits: usize,
shift: F,
transform: T,
) -> Self::Evaluations
where
T: FnOnce(&mut RowMajorMatrixViewMut<'_, F>, Layout),
{
let mut coeffs = self.idft_batch(mat);
transform(&mut coeffs.as_view_mut(), Layout::Natural);
coeffs.values.resize(
coeffs
.values
.len()
.checked_shl(added_bits.try_into().unwrap())
.unwrap(),
F::ZERO,
);
self.coset_dft_batch(coeffs, shift)
}
fn dft_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(&self, vec: Vec<V>) -> Vec<V> {
self.dft_algebra_batch(RowMajorMatrix::new_col(vec)).values
}
fn dft_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self.dft_batch(base_mat).to_row_major_matrix();
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
fn coset_dft_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
vec: Vec<V>,
shift: F,
) -> Vec<V> {
self.coset_dft_algebra_batch(RowMajorMatrix::new_col(vec), shift)
.to_row_major_matrix()
.values
}
fn coset_dft_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
shift: F,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self.coset_dft_batch(base_mat, shift).to_row_major_matrix();
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
fn idft_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(&self, vec: Vec<V>) -> Vec<V> {
self.idft_algebra_batch(RowMajorMatrix::new_col(vec)).values
}
fn idft_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self.idft_batch(base_mat);
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
fn coset_idft_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
vec: Vec<V>,
shift: F,
) -> Vec<V> {
self.coset_idft_algebra_batch(RowMajorMatrix::new_col(vec), shift)
.values
}
fn coset_idft_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
shift: F,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self.coset_idft_batch(base_mat, shift);
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
fn lde_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
vec: Vec<V>,
added_bits: usize,
) -> Vec<V> {
self.lde_algebra_batch(RowMajorMatrix::new_col(vec), added_bits)
.to_row_major_matrix()
.values
}
fn lde_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
added_bits: usize,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self.lde_batch(base_mat, added_bits).to_row_major_matrix();
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
fn coset_lde_algebra<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
vec: Vec<V>,
added_bits: usize,
shift: F,
) -> Vec<V> {
self.coset_lde_algebra_batch(RowMajorMatrix::new_col(vec), added_bits, shift)
.to_row_major_matrix()
.values
}
fn coset_lde_algebra_batch<V: BasedVectorSpace<F> + Clone + Send + Sync>(
&self,
mat: RowMajorMatrix<V>,
added_bits: usize,
shift: F,
) -> RowMajorMatrix<V> {
let init_width = mat.width();
let base_mat =
RowMajorMatrix::new(V::flatten_to_base(mat.values), init_width * V::DIMENSION);
let base_dft_output = self
.coset_lde_batch(base_mat, added_bits, shift)
.to_row_major_matrix();
RowMajorMatrix::new(
V::reconstitute_from_base(base_dft_output.values),
init_width,
)
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum Layout {
Natural,
BitReversed,
}