use p3_dft::TwoAdicSubgroupDft;
use p3_field::{Field, TwoAdicField};
use p3_matrix::Matrix;
use p3_matrix::dense::RowMajorMatrix;
pub trait Encoder<F: Field> {
fn encode_batch(&self, message: RowMajorMatrix<F>, log_inv_rate: usize) -> RowMajorMatrix<F>;
}
impl<F: TwoAdicField, D: TwoAdicSubgroupDft<F>> Encoder<F> for D {
fn encode_batch(
&self,
mut message: RowMajorMatrix<F>,
log_inv_rate: usize,
) -> RowMajorMatrix<F> {
if log_inv_rate > 0 {
let len = message.values.len();
let padded_len = u32::try_from(log_inv_rate)
.ok()
.and_then(|rate| len.checked_shl(rate))
.filter(|&padded| padded >> log_inv_rate == len)
.expect("codeword length overflows usize");
let mut values = F::zero_vec(padded_len);
values[..len].copy_from_slice(&message.values);
message.values = values;
}
self.dft_batch(message).to_row_major_matrix()
}
}
#[cfg(test)]
mod tests {
use alloc::vec;
use p3_baby_bear::BabyBear;
use p3_dft::{Radix2DFTSmallBatch, Radix2DitParallel, TwoAdicSubgroupDft};
use p3_field::PrimeCharacteristicRing;
use p3_matrix::Matrix;
use p3_matrix::dense::RowMajorMatrix;
use rand::SeedableRng;
use rand::rngs::SmallRng;
use super::Encoder;
fn check_matches_padded_dft<D: TwoAdicSubgroupDft<BabyBear>>(dft: &D) {
let mut rng = SmallRng::seed_from_u64(1);
let message = RowMajorMatrix::<BabyBear>::rand(&mut rng, 8, 3);
let mut padded = message.clone();
padded
.values
.resize(message.values.len() * 4, BabyBear::ZERO);
let expected = dft.dft_batch(padded).to_row_major_matrix();
assert_eq!(dft.encode_batch(message, 2), expected);
}
#[test]
fn small_batch_encoder_matches_padded_dft() {
check_matches_padded_dft(&Radix2DFTSmallBatch::<BabyBear>::default());
}
#[test]
fn dit_parallel_encoder_matches_padded_dft() {
check_matches_padded_dft(&Radix2DitParallel::<BabyBear>::default());
}
#[test]
#[should_panic = "codeword length overflows usize"]
fn encode_batch_panics_when_the_codeword_length_overflows() {
let message = RowMajorMatrix::<BabyBear>::new(vec![BabyBear::ZERO; 2], 1);
let _ = Radix2DitParallel::<BabyBear>::default()
.encode_batch(message, usize::BITS as usize - 1);
}
}