use crate::air::Assertion;
use core::fmt::{Display, Formatter};
use math::{log2, FieldElement, StarkField};
use utils::collections::Vec;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConstraintDivisor<B: StarkField> {
pub(super) numerator: Vec<(usize, B)>,
pub(super) exemptions: Vec<B>,
}
impl<B: StarkField> ConstraintDivisor<B> {
fn new(numerator: Vec<(usize, B)>, exemptions: Vec<B>) -> Self {
ConstraintDivisor {
numerator,
exemptions,
}
}
pub fn from_transition(trace_length: usize, num_exemptions: usize) -> Self {
assert!(
num_exemptions > 0,
"invalid number of transition exemptions: must be greater than zero"
);
let exemptions = (trace_length - num_exemptions..trace_length)
.map(|step| get_trace_domain_value_at::<B>(trace_length, step))
.collect();
Self::new(vec![(trace_length, B::ONE)], exemptions)
}
pub fn from_assertion<E>(assertion: &Assertion<E>, trace_length: usize) -> Self
where
E: FieldElement<BaseField = B>,
{
let num_steps = assertion.get_num_steps(trace_length);
if assertion.first_step == 0 {
Self::new(vec![(num_steps, B::ONE)], vec![])
} else {
let trace_offset = num_steps * assertion.first_step;
let offset = get_trace_domain_value_at::<B>(trace_length, trace_offset);
Self::new(vec![(num_steps, offset)], vec![])
}
}
pub fn numerator(&self) -> &[(usize, B)] {
&self.numerator
}
pub fn exemptions(&self) -> &[B] {
&self.exemptions
}
pub fn degree(&self) -> usize {
let numerator_degree = self
.numerator
.iter()
.fold(0, |degree, term| degree + term.0);
let denominator_degree = self.exemptions.len();
numerator_degree - denominator_degree
}
pub fn evaluate_at<E: FieldElement<BaseField = B>>(&self, x: E) -> E {
let mut numerator = E::ONE;
for (degree, constant) in self.numerator.iter() {
let v = x.exp((*degree as u32).into());
let v = v - E::from(*constant);
numerator *= v;
}
let denominator = self.evaluate_exemptions_at(x);
numerator / denominator
}
#[inline(always)]
pub fn evaluate_exemptions_at<E: FieldElement<BaseField = B>>(&self, x: E) -> E {
self.exemptions
.iter()
.fold(E::ONE, |r, &e| r * (x - E::from(e)))
}
}
impl<B: StarkField> Display for ConstraintDivisor<B> {
fn fmt(&self, f: &mut Formatter) -> core::fmt::Result {
for (degree, offset) in self.numerator.iter() {
write!(f, "(x^{} - {})", degree, offset)?;
}
if !self.exemptions.is_empty() {
write!(f, " / ")?;
for x in self.exemptions.iter() {
write!(f, "(x - {})", x)?;
}
}
Ok(())
}
}
fn get_trace_domain_value_at<B: StarkField>(trace_length: usize, step: usize) -> B {
debug_assert!(
step < trace_length,
"step must be in the trace domain [0, {})",
trace_length
);
let g = B::get_root_of_unity(log2(trace_length));
g.exp((step as u64).into())
}
#[cfg(test)]
mod tests {
use super::*;
use math::{fields::f128::BaseElement, polynom};
#[test]
fn constraint_divisor_degree() {
let div = ConstraintDivisor::new(vec![(4, BaseElement::ONE)], vec![]);
assert_eq!(4, div.degree());
let div = ConstraintDivisor::new(
vec![
(4, BaseElement::ONE),
(2, BaseElement::new(2)),
(3, BaseElement::new(3)),
],
vec![],
);
assert_eq!(9, div.degree());
let div = ConstraintDivisor::new(
vec![
(4, BaseElement::ONE),
(2, BaseElement::new(2)),
(3, BaseElement::new(3)),
],
vec![BaseElement::ONE, BaseElement::new(2)],
);
assert_eq!(7, div.degree());
}
#[test]
fn constraint_divisor_evaluation() {
let div = ConstraintDivisor::new(vec![(4, BaseElement::ONE)], vec![]);
assert_eq!(BaseElement::new(15), div.evaluate_at(BaseElement::new(2)));
let div = ConstraintDivisor::new(
vec![
(4, BaseElement::ONE),
(2, BaseElement::new(2)),
(3, BaseElement::new(3)),
],
vec![],
);
let expected = BaseElement::new(15) * BaseElement::new(2) * BaseElement::new(5);
assert_eq!(expected, div.evaluate_at(BaseElement::new(2)));
let div = ConstraintDivisor::new(
vec![
(4, BaseElement::ONE),
(2, BaseElement::new(2)),
(3, BaseElement::new(3)),
],
vec![BaseElement::ONE, BaseElement::new(2)],
);
let expected = BaseElement::new(255) * BaseElement::new(14) * BaseElement::new(61)
/ BaseElement::new(6);
assert_eq!(expected, div.evaluate_at(BaseElement::new(4)));
}
#[test]
fn constraint_divisor_equivalence() {
let n = 8_usize;
let g = BaseElement::get_root_of_unity(n.trailing_zeros());
let k = 4 as u32;
let j = n as u32 / k;
let assertion = Assertion::periodic(0, 0, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
let poly = polynom::mul(
&polynom::mul(
&[-BaseElement::ONE, BaseElement::ONE],
&[-g.exp(j.into()), BaseElement::ONE],
),
&polynom::mul(
&[-g.exp((2 * j).into()), BaseElement::ONE],
&[-g.exp((3 * j).into()), BaseElement::ONE],
),
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i % (j as usize) == 0 {
assert_eq!(BaseElement::ZERO, actual);
}
}
let offset = 1u32;
let assertion = Assertion::periodic(0, offset as usize, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
assert_eq!(
ConstraintDivisor::new(vec![(k as usize, g.exp(k.into()))], vec![]),
divisor
);
let poly = polynom::mul(
&polynom::mul(
&[-g.exp(offset.into()), BaseElement::ONE],
&[-g.exp((offset + j).into()), BaseElement::ONE],
),
&polynom::mul(
&[-g.exp((offset + 2 * j).into()), BaseElement::ONE],
&[-g.exp((offset + 3 * j).into()), BaseElement::ONE],
),
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i % (j as usize) == offset as usize {
assert_eq!(BaseElement::ZERO, actual);
}
}
let offset = 3u32;
let k = 2 as u32;
let j = n as u32 / k;
let assertion = Assertion::periodic(0, offset as usize, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
assert_eq!(
ConstraintDivisor::new(vec![(k as usize, g.exp((offset * k).into()))], vec![]),
divisor
);
let poly = polynom::mul(
&[-g.exp(offset.into()), BaseElement::ONE],
&[-g.exp((offset + j).into()), BaseElement::ONE],
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i % (j as usize) == offset as usize {
assert_eq!(BaseElement::ZERO, actual);
}
}
}
}