use alloc::vec::Vec;
use itertools::Itertools;
use lib_q_stark_field::coset::TwoAdicMultiplicativeCoset;
use lib_q_stark_field::{
ExtensionField,
Field,
TwoAdicField,
batch_multiplicative_inverse,
};
use lib_q_stark_matrix::Matrix;
use lib_q_stark_matrix::dense::RowMajorMatrix;
use lib_q_stark_util::{
log2_ceil_usize,
log2_strict_usize,
};
#[derive(Debug)]
pub struct LagrangeSelectors<T> {
pub is_first_row: T,
pub is_last_row: T,
pub is_transition: T,
pub inv_vanishing: T,
}
pub trait PolynomialSpace: Copy {
type Val: Field;
fn size(&self) -> usize;
fn first_point(&self) -> Self::Val;
fn next_point<Ext: ExtensionField<Self::Val>>(&self, x: Ext) -> Option<Ext>;
fn create_disjoint_domain(&self, min_size: usize) -> Self;
fn try_create_disjoint_domain(&self, min_size: usize) -> Option<Self> {
Some(self.create_disjoint_domain(min_size))
}
fn split_domains(&self, num_chunks: usize) -> Vec<Self>;
fn split_evals(
&self,
num_chunks: usize,
evals: RowMajorMatrix<Self::Val>,
) -> Vec<RowMajorMatrix<Self::Val>>;
fn vanishing_poly_at_point<Ext: ExtensionField<Self::Val>>(&self, point: Ext) -> Ext;
fn selectors_at_point<Ext: ExtensionField<Self::Val>>(
&self,
point: Ext,
) -> LagrangeSelectors<Ext>;
fn selectors_on_coset(&self, coset: Self) -> LagrangeSelectors<Vec<Self::Val>>;
}
impl<Val: TwoAdicField> PolynomialSpace for TwoAdicMultiplicativeCoset<Val> {
type Val = Val;
fn size(&self) -> usize {
self.size()
}
fn first_point(&self) -> Self::Val {
self.shift()
}
fn next_point<Ext: ExtensionField<Val>>(&self, x: Ext) -> Option<Ext> {
Some(x * self.subgroup_generator())
}
fn create_disjoint_domain(&self, min_size: usize) -> Self {
self.try_create_disjoint_domain(min_size)
.unwrap_or_else(|| {
panic!(
"create_disjoint_domain: min_size {min_size} exceeds 1 << Val::TWO_ADICITY \
({}); use try_create_disjoint_domain to handle this without panicking",
Val::TWO_ADICITY
)
})
}
fn try_create_disjoint_domain(&self, min_size: usize) -> Option<Self> {
Self::new(self.shift() * Val::GENERATOR, log2_ceil_usize(min_size))
}
fn split_domains(&self, num_chunks: usize) -> Vec<Self> {
let log_chunks = log2_strict_usize(num_chunks);
debug_assert!(log_chunks <= self.log_size());
(0..num_chunks)
.map(|i| {
Self::new(
self.shift() * self.subgroup_generator().exp_u64(i as u64),
self.log_size() - log_chunks,
)
.unwrap() })
.collect()
}
fn split_evals(
&self,
num_chunks: usize,
evals: RowMajorMatrix<Self::Val>,
) -> Vec<RowMajorMatrix<Self::Val>> {
debug_assert_eq!(evals.height(), self.size());
debug_assert!(log2_strict_usize(num_chunks) <= self.log_size());
let height = evals.height();
let width = evals.width();
let rows_per_chunk = height / num_chunks;
let mut values: Vec<Vec<Self::Val>> = (0..num_chunks)
.map(|_| Self::Val::zero_vec(rows_per_chunk * width))
.collect();
for i in 0..rows_per_chunk {
let base_row = i * num_chunks;
let dst_start = i * width;
let dst_end = dst_start + width;
for (chunk, dst_vec) in values.iter_mut().enumerate().take(num_chunks) {
let r = base_row + chunk;
let row = unsafe { evals.row_slice_unchecked(r) };
dst_vec[dst_start..dst_end].copy_from_slice(&row);
}
}
values
.into_iter()
.map(|v| RowMajorMatrix::new(v, width))
.collect()
}
fn vanishing_poly_at_point<Ext: ExtensionField<Val>>(&self, point: Ext) -> Ext {
(point * self.shift_inverse()).exp_power_of_2(self.log_size()) - Ext::ONE
}
fn selectors_at_point<Ext: ExtensionField<Val>>(&self, point: Ext) -> LagrangeSelectors<Ext> {
let unshifted_point = point * self.shift_inverse();
let z_h = unshifted_point.exp_power_of_2(self.log_size()) - Ext::ONE;
LagrangeSelectors {
is_first_row: z_h / (unshifted_point - Ext::ONE),
is_last_row: z_h / (unshifted_point - self.subgroup_generator().inverse()),
is_transition: unshifted_point - self.subgroup_generator().inverse(),
inv_vanishing: z_h.inverse(),
}
}
fn selectors_on_coset(&self, coset: Self) -> LagrangeSelectors<Vec<Val>> {
assert_eq!(self.shift(), Val::ONE);
assert_ne!(coset.shift(), Val::ONE);
assert!(coset.log_size() >= self.log_size());
let rate_bits = coset.log_size() - self.log_size();
let s_pow_n = coset.shift().exp_power_of_2(self.log_size());
let evals = Val::two_adic_generator(rate_bits)
.powers()
.take(1 << rate_bits)
.map(|x| s_pow_n * x - Val::ONE)
.collect_vec();
let xs = coset.iter().collect();
let single_point_selector = |i: u64| {
let coset_i = self.subgroup_generator().exp_u64(i);
let denoms = xs.iter().map(|&x| x - coset_i).collect_vec();
let invs = batch_multiplicative_inverse(&denoms);
evals
.iter()
.cycle()
.zip(invs)
.map(|(&z_h, inv)| z_h * inv)
.collect_vec()
};
let subgroup_last = self.subgroup_generator().inverse();
LagrangeSelectors {
is_first_row: single_point_selector(0),
is_last_row: single_point_selector(self.size() as u64 - 1),
is_transition: xs.into_iter().map(|x| x - subgroup_last).collect(),
inv_vanishing: batch_multiplicative_inverse(&evals)
.into_iter()
.cycle()
.take(coset.size())
.collect(),
}
}
}
#[cfg(test)]
mod tests {
use lib_q_stark_baby_bear::BabyBear;
use lib_q_stark_field::{
PrimeCharacteristicRing,
PrimeField32,
};
use super::*;
type F = BabyBear;
fn coset(shift: F, log_size: usize) -> TwoAdicMultiplicativeCoset<F> {
TwoAdicMultiplicativeCoset::new(shift, log_size).unwrap()
}
fn sorted_u32(points: impl IntoIterator<Item = F>) -> Vec<u32> {
let mut v: Vec<u32> = points.into_iter().map(|p| p.as_canonical_u32()).collect();
v.sort_unstable();
v
}
#[test]
fn size_and_first_point_match_constructor_arguments() {
let c = coset(F::new(7), 3);
assert_eq!(PolynomialSpace::size(&c), 8);
assert_eq!(PolynomialSpace::first_point(&c), F::new(7));
}
#[test]
fn next_point_is_multiplication_by_the_subgroup_generator() {
let c = coset(F::new(5), 4);
let g = c.subgroup_generator();
let x = F::new(123);
assert_eq!(PolynomialSpace::next_point::<F>(&c, x), Some(x * g));
}
#[test]
fn next_point_returns_to_start_after_exactly_size_steps() {
let c = coset(F::new(11), 5);
let mut x = PolynomialSpace::first_point(&c);
for _ in 0..PolynomialSpace::size(&c) {
x = PolynomialSpace::next_point::<F>(&c, x).unwrap();
}
assert_eq!(x, PolynomialSpace::first_point(&c));
}
#[test]
fn vanishing_poly_is_zero_exactly_on_the_coset() {
let c = coset(F::new(9), 4);
for point in c.iter() {
assert_eq!(
PolynomialSpace::vanishing_poly_at_point::<F>(&c, point),
F::ZERO
);
}
let disjoint = PolynomialSpace::create_disjoint_domain(&c, PolynomialSpace::size(&c));
for point in disjoint.iter() {
assert_ne!(
PolynomialSpace::vanishing_poly_at_point::<F>(&c, point),
F::ZERO
);
}
}
#[test]
fn vanishing_poly_negative_control_distinguishes_in_from_out_of_domain() {
let c = coset(F::new(9), 4);
let on_domain_is_zero =
PolynomialSpace::vanishing_poly_at_point::<F>(&c, PolynomialSpace::first_point(&c)) ==
F::ZERO;
let disjoint = PolynomialSpace::create_disjoint_domain(&c, PolynomialSpace::size(&c));
let off_domain_is_zero = PolynomialSpace::vanishing_poly_at_point::<F>(
&c,
PolynomialSpace::first_point(&disjoint),
) == F::ZERO;
assert_ne!(on_domain_is_zero, off_domain_is_zero);
}
#[test]
fn create_disjoint_domain_has_correct_size_and_shares_no_point() {
let c = coset(F::new(13), 3); let c_points = sorted_u32(c.iter());
for min_size in [1usize, 2, 3, 7, 8, 9, 20, 64] {
let k = PolynomialSpace::create_disjoint_domain(&c, min_size);
assert_eq!(PolynomialSpace::size(&k), min_size.next_power_of_two());
assert_eq!(k.shift(), c.shift() * F::GENERATOR);
let k_points = sorted_u32(k.iter());
let mut merged = c_points.clone();
merged.extend(&k_points);
merged.sort_unstable();
merged.dedup();
assert_eq!(merged.len(), c_points.len() + k_points.len());
}
}
#[test]
fn try_create_disjoint_domain_agrees_with_the_infallible_version_in_range() {
let c = coset(F::new(13), 3);
for min_size in [1usize, 5, 8, 100] {
assert_eq!(
PolynomialSpace::try_create_disjoint_domain(&c, min_size).map(|k| k.shift()),
Some(PolynomialSpace::create_disjoint_domain(&c, min_size).shift())
);
}
}
#[test]
fn try_create_disjoint_domain_rejects_out_of_range_min_size() {
let c = coset(F::new(13), 3);
assert!(PolynomialSpace::try_create_disjoint_domain(&c, 1 << F::TWO_ADICITY).is_some());
assert!(
PolynomialSpace::try_create_disjoint_domain(&c, (1 << F::TWO_ADICITY) + 1).is_none()
);
}
#[test]
#[should_panic(expected = "exceeds 1 << Val::TWO_ADICITY")]
fn create_disjoint_domain_panics_out_of_range_where_try_returns_none() {
let c = coset(F::new(13), 3);
let _ = PolynomialSpace::create_disjoint_domain(&c, (1 << F::TWO_ADICITY) + 1);
}
#[test]
fn split_domains_partition_the_original_coset() {
let c = coset(F::new(17), 4); let orig_points = sorted_u32(c.iter());
for &num_chunks in &[1usize, 2, 4, 8, 16] {
let subs = PolynomialSpace::split_domains(&c, num_chunks);
assert_eq!(subs.len(), num_chunks);
let mut all_points: Vec<u32> = Vec::new();
for s in &subs {
assert_eq!(PolynomialSpace::size(s), c.size() / num_chunks);
all_points.extend(s.iter().map(|p| p.as_canonical_u32()));
}
all_points.sort_unstable();
assert_eq!(
all_points, orig_points,
"split_domains({num_chunks}) did not exactly partition the original coset"
);
}
}
#[test]
fn split_domains_negative_control_detects_a_missing_chunk() {
let c = coset(F::new(17), 4);
let orig_points = sorted_u32(c.iter());
let subs = PolynomialSpace::split_domains(&c, 4);
let mut all_points: Vec<u32> = Vec::new();
for s in &subs[..subs.len() - 1] {
all_points.extend(s.iter().map(|p| p.as_canonical_u32()));
}
all_points.sort_unstable();
assert_ne!(all_points, orig_points);
}
#[test]
fn split_evals_decimates_rows_to_match_split_domains() {
let c = coset(F::new(3), 4); let width = 2;
let height = c.size();
let values: Vec<F> = (0..height * width)
.map(|i| F::new((i / width) as u32))
.collect();
let evals = RowMajorMatrix::new(values, width);
let num_chunks = 4;
let chunks = PolynomialSpace::split_evals(&c, num_chunks, evals);
assert_eq!(chunks.len(), num_chunks);
let rows_per_chunk = height / num_chunks;
for (chunk_idx, chunk) in chunks.iter().enumerate() {
assert_eq!(chunk.height(), rows_per_chunk);
for i in 0..rows_per_chunk {
let expected_row = i * num_chunks + chunk_idx;
let row = chunk.row_slice(i).unwrap();
assert_eq!(row[0], F::new(expected_row as u32));
}
}
}
#[test]
fn selectors_on_coset_agrees_with_selectors_at_point_for_every_point() {
let h = coset(F::ONE, 3); let disjoint_coset = PolynomialSpace::create_disjoint_domain(&h, 2 * h.size());
let batched = PolynomialSpace::selectors_on_coset(&h, disjoint_coset);
let points: Vec<F> = disjoint_coset.iter().collect();
assert_eq!(points.len(), batched.is_first_row.len());
for (i, &x) in points.iter().enumerate() {
let single = PolynomialSpace::selectors_at_point::<F>(&h, x);
assert_eq!(single.is_first_row, batched.is_first_row[i]);
assert_eq!(single.is_last_row, batched.is_last_row[i]);
assert_eq!(single.is_transition, batched.is_transition[i]);
assert_eq!(single.inv_vanishing, batched.inv_vanishing[i]);
}
}
#[test]
fn selectors_negative_control_values_actually_vary_by_point() {
let h = coset(F::ONE, 3);
let disjoint_coset = PolynomialSpace::create_disjoint_domain(&h, 2 * h.size());
let mut points = disjoint_coset.iter();
let x0 = points.next().unwrap();
let x1 = points.next().unwrap();
let s0 = PolynomialSpace::selectors_at_point::<F>(&h, x0);
let s1 = PolynomialSpace::selectors_at_point::<F>(&h, x1);
assert_ne!(s0.is_first_row, s1.is_first_row);
}
}