use crate::{
CommitmentKey,
bellpepper::{
r1cs::{MultiRoundSpartanWitness, MultiRoundState},
solver::SatisfyingAssignment,
},
big_num::DelayedReduction,
errors::SpartanError,
polys::{
multilinear::MultilinearPolynomial,
univariate::{CompressedUniPoly, UniPoly},
},
r1cs::SplitMultiRoundR1CSShape,
start_span,
traits::{Engine, transcript::TranscriptEngineTrait},
zk::{NeutronNovaVerifierCircuit, SpartanVerifierCircuit},
};
use ff::Field;
use num_traits::Zero;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use tracing::info;
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(bound = "")]
pub struct SumcheckProof<E: Engine> {
compressed_polys: Vec<CompressedUniPoly<E::Scalar>>,
}
impl<E: Engine> SumcheckProof<E> {
pub fn verify(
&self,
claim: E::Scalar,
num_rounds: usize,
degree_bound: usize,
transcript: &mut E::TE,
) -> Result<(E::Scalar, Vec<E::Scalar>), SpartanError> {
let (_verify_span, verify_t) = start_span!("sumcheck_verify");
let mut e = claim;
let mut r: Vec<E::Scalar> = Vec::new();
if self.compressed_polys.len() != num_rounds {
return Err(SpartanError::InvalidSumcheckProof);
}
for i in 0..self.compressed_polys.len() {
let (_round_span, round_t) = start_span!("sumcheck_verify_round", round = i);
let poly = self.compressed_polys[i].decompress(&e);
if poly.degree() != degree_bound {
return Err(SpartanError::InvalidSumcheckProof);
}
debug_assert_eq!(poly.eval_at_zero() + poly.eval_at_one(), e);
transcript.absorb(b"p", &poly);
let r_i = transcript.squeeze(b"c")?;
r.push(r_i);
e = poly.evaluate(&r_i);
if round_t.elapsed().as_millis() > 0 {
info!(elapsed_ms = %round_t.elapsed().as_millis(), "sumcheck_verify_round");
}
}
info!(elapsed_ms = %verify_t.elapsed().as_millis(), "sumcheck_verify");
Ok((e, r))
}
#[inline]
fn compute_eval_points_quad(
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
) -> (E::Scalar, E::Scalar) {
type Acc<S> = <S as DelayedReduction<S>>::Accumulator;
let len = poly_A.Z.len() / 2;
let (acc_0, acc_2) = (0..len)
.into_par_iter()
.fold(
|| (Acc::<E::Scalar>::zero(), Acc::<E::Scalar>::zero()),
|mut acc, i| {
let a_low = &poly_A[i];
let a_high = &poly_A[len + i];
let b_low = &poly_B[i];
let b_high = &poly_B[len + i];
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.0, a_low, b_low,
);
let a_bound = *a_high + *a_high - *a_low;
let b_bound = *b_high + *b_high - *b_low;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.1, &a_bound, &b_bound,
);
acc
},
)
.reduce(
|| (Acc::<E::Scalar>::zero(), Acc::<E::Scalar>::zero()),
|mut a, b| {
a.0 += b.0;
a.1 += b.1;
a
},
);
(
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_0),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_2),
)
}
pub fn prove_quad(
claim: &E::Scalar,
num_rounds: usize,
poly_A: &mut MultilinearPolynomial<E::Scalar>,
poly_B: &mut MultilinearPolynomial<E::Scalar>,
transcript: &mut E::TE,
) -> Result<(Self, Vec<E::Scalar>, Vec<E::Scalar>), SpartanError> {
let mut r: Vec<E::Scalar> = Vec::new();
let mut polys: Vec<CompressedUniPoly<E::Scalar>> = Vec::new();
let mut claim_per_round = *claim;
for round in 0..num_rounds {
let (_round_span, round_t) = start_span!("sumcheck_quad_round", round = round);
let poly = {
let (_eval_span, eval_t) = start_span!("compute_eval_points_quad");
let (eval_point_0, eval_point_2) = Self::compute_eval_points_quad(poly_A, poly_B);
if eval_t.elapsed().as_millis() > 0 {
info!(elapsed_ms = %eval_t.elapsed().as_millis(), "compute_eval_points_quad");
}
let evals = vec![eval_point_0, claim_per_round - eval_point_0, eval_point_2];
UniPoly::from_evals(&evals)?
};
transcript.absorb(b"p", &poly);
let r_i = transcript.squeeze(b"c")?;
r.push(r_i);
polys.push(poly.compress());
claim_per_round = poly.evaluate(&r_i);
let (_bind_span, bind_t) = start_span!("bind_poly_vars_quad");
rayon::join(
|| poly_A.bind_poly_var_top(&r_i),
|| poly_B.bind_poly_var_top(&r_i),
);
info!(elapsed_ms = %bind_t.elapsed().as_millis(), "bind_poly_vars_quad");
info!(elapsed_ms = %round_t.elapsed().as_millis(), round = round, "sumcheck_quad_round");
}
Ok((
SumcheckProof {
compressed_polys: polys,
},
r,
vec![poly_A[0], poly_B[0]],
))
}
#[inline]
fn compute_eval_points_cubic_with_additive_term(
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
poly_C: &MultilinearPolynomial<E::Scalar>,
poly_D: &MultilinearPolynomial<E::Scalar>,
) -> (E::Scalar, E::Scalar, E::Scalar) {
type Acc<S> = <S as DelayedReduction<S>>::Accumulator;
let len = poly_B.Z.len() / 2;
let (acc_0, acc_2, acc_3) = (0..len)
.into_par_iter()
.fold(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut acc, i| {
let a_low = &poly_A[i];
let a_high = &poly_A[i + len];
let b_low = &poly_B[i];
let b_high = &poly_B[i + len];
let c_low = &poly_C[i];
let c_high = &poly_C[i + len];
let d_low = &poly_D[i];
let d_high = &poly_D[i + len];
let inner_0 = *b_low * *c_low - *d_low;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.0, a_low, &inner_0,
);
let a_bound = *a_high + *a_high - *a_low;
let b_bound = *b_high + *b_high - *b_low;
let c_bound = *c_high + *c_high - *c_low;
let d_bound = *d_high + *d_high - *d_low;
let inner_2 = b_bound * c_bound - d_bound;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.1, &a_bound, &inner_2,
);
let a_bound = a_bound + *a_high - *a_low;
let b_bound = b_bound + *b_high - *b_low;
let c_bound = c_bound + *c_high - *c_low;
let d_bound = d_bound + *d_high - *d_low;
let inner_3 = b_bound * c_bound - d_bound;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.2, &a_bound, &inner_3,
);
acc
},
)
.reduce(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut a, b| {
a.0 += b.0;
a.1 += b.1;
a.2 += b.2;
a
},
);
(
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_0),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_2),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_3),
)
}
#[inline]
fn compute_eval_points_cubic_with_additive_term_with_outer_pow(
pow_tau_left: &MultilinearPolynomial<E::Scalar>,
pow_tau_right: &MultilinearPolynomial<E::Scalar>,
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
poly_C: &MultilinearPolynomial<E::Scalar>,
) -> (E::Scalar, E::Scalar, E::Scalar) {
type Acc<S> = <S as DelayedReduction<S>>::Accumulator;
let len = poly_A.Z.len() / 2;
let left = pow_tau_left.Z.len();
if len < left {
return Self::compute_eval_points_cubic_with_additive_term(
pow_tau_left,
poly_A,
poly_B,
poly_C,
);
}
let right = len / left;
let (acc_0, acc_2, acc_3) = (0..left)
.into_par_iter()
.fold(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut outer_acc, i| {
let pow_left = &pow_tau_left[i];
let mut inner_0 = Acc::<E::Scalar>::zero();
let mut inner_2 = Acc::<E::Scalar>::zero();
let mut inner_3 = Acc::<E::Scalar>::zero();
for j in 0..right {
let low = i + j * left;
let high = low + len;
let tau_low = &pow_tau_right[j];
let tau_high = &pow_tau_right[j + right];
let a_low = &poly_A[low];
let a_high = &poly_A[high];
let b_low = &poly_B[low];
let b_high = &poly_B[high];
let c_low = &poly_C[low];
let c_high = &poly_C[high];
let prod_0 = *a_low * *b_low - *c_low;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_0,
tau_low,
&prod_0,
);
let tau_bound = *tau_high + *tau_high - *tau_low;
let a_bound = *a_high + *a_high - *a_low;
let b_bound = *b_high + *b_high - *b_low;
let c_bound = *c_high + *c_high - *c_low;
let prod_2 = a_bound * b_bound - c_bound;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_2,
&tau_bound,
&prod_2,
);
let tau_bound = tau_bound + *tau_high - *tau_low;
let a_bound = a_bound + *a_high - *a_low;
let b_bound = b_bound + *b_high - *b_low;
let c_bound = c_bound + *c_high - *c_low;
let prod_3 = a_bound * b_bound - c_bound;
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_3,
&tau_bound,
&prod_3,
);
}
let inner_0_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_0);
let inner_2_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_2);
let inner_3_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_3);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.0,
pow_left,
&inner_0_red,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.1,
pow_left,
&inner_2_red,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.2,
pow_left,
&inner_3_red,
);
outer_acc
},
)
.reduce(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut a, b| {
a.0 += b.0;
a.1 += b.1;
a.2 += b.2;
a
},
);
(
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_0),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_2),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_3),
)
}
pub fn prove_cubic_with_three_inputs(
claim: &E::Scalar,
taus: Vec<E::Scalar>,
poly_A: &mut MultilinearPolynomial<E::Scalar>,
poly_B: &mut MultilinearPolynomial<E::Scalar>,
poly_C: &mut MultilinearPolynomial<E::Scalar>,
transcript: &mut E::TE,
) -> Result<(Self, Vec<E::Scalar>, Vec<E::Scalar>), SpartanError> {
let mut r: Vec<E::Scalar> = Vec::new();
let mut polys: Vec<CompressedUniPoly<E::Scalar>> = Vec::new();
let mut claim_per_round = *claim;
let num_rounds = taus.len();
let mut eq_instance = eq_sumcheck::EqSumCheckInstance::<E>::new(taus);
for round in 0..num_rounds {
let (_round_span, round_t) = start_span!("sumcheck_round", round = round);
let poly = {
let (_eval_span, eval_t) = start_span!("compute_eval_points");
let (eval_point_0, eval_point_2, eval_point_3) =
eq_instance.evaluation_points_cubic_with_three_inputs(round, poly_A, poly_B, poly_C);
if eval_t.elapsed().as_millis() > 0 {
info!(elapsed_ms = %eval_t.elapsed().as_millis(), "compute_eval_points");
}
let evals = vec![
eval_point_0,
claim_per_round - eval_point_0,
eval_point_2,
eval_point_3,
];
UniPoly::from_evals(&evals)?
};
transcript.absorb(b"p", &poly);
let r_i = transcript.squeeze(b"c")?;
r.push(r_i);
polys.push(poly.compress());
claim_per_round = poly.evaluate(&r_i);
let (_bind_span, bind_t) = start_span!("bind_poly_vars");
rayon::join(
|| poly_A.bind_poly_var_top(&r_i),
|| poly_B.bind_poly_var_top(&r_i),
);
rayon::join(
|| poly_C.bind_poly_var_top(&r_i),
|| eq_instance.bound(&r_i),
);
info!(elapsed_ms = %bind_t.elapsed().as_millis(), "bind_poly_vars");
info!(elapsed_ms = %round_t.elapsed().as_millis(), round = round, "sumcheck_round");
}
Ok((
SumcheckProof {
compressed_polys: polys,
},
r,
vec![poly_A[0], poly_B[0], poly_C[0]],
))
}
pub fn prove_cubic_with_additive_term_zk(
num_rounds: usize,
taus: &[E::Scalar],
poly_Az: &mut MultilinearPolynomial<E::Scalar>,
poly_Bz: &mut MultilinearPolynomial<E::Scalar>,
poly_Cz: &mut MultilinearPolynomial<E::Scalar>,
verifier_circuit: &mut SpartanVerifierCircuit<E>,
state: &mut MultiRoundState<E>,
vc_shape: &SplitMultiRoundR1CSShape<E>,
vc_ck: &CommitmentKey<E>,
transcript: &mut E::TE,
) -> Result<Vec<E::Scalar>, SpartanError> {
let mut r_x: Vec<E::Scalar> = Vec::with_capacity(num_rounds);
let mut claim_outer_round = E::Scalar::ZERO;
let mut eq_instance = eq_sumcheck::EqSumCheckInstance::<E>::new(taus.to_vec());
for i in 0..num_rounds {
let (eval0, eval2, eval3) =
eq_instance.evaluation_points_cubic_with_three_inputs(i, poly_Az, poly_Bz, poly_Cz);
let evals = vec![eval0, claim_outer_round - eval0, eval2, eval3];
let poly = UniPoly::from_evals(&evals)?;
verifier_circuit.outer_polys[i] = [
poly.coeffs[0],
poly.coeffs[1],
poly.coeffs[2],
poly.coeffs[3],
];
let chals = SatisfyingAssignment::<E>::process_round(
state,
vc_shape,
vc_ck,
verifier_circuit,
i,
transcript,
)?;
r_x.push(chals[0]);
claim_outer_round = poly.evaluate(&chals[0]);
rayon::join(
|| poly_Az.bind_poly_var_top(&chals[0]),
|| {
rayon::join(
|| poly_Bz.bind_poly_var_top(&chals[0]),
|| poly_Cz.bind_poly_var_top(&chals[0]),
);
},
);
eq_instance.bound(&chals[0]);
}
Ok(r_x)
}
pub fn prove_quad_zk(
claim: &E::Scalar,
num_rounds: usize,
poly_ABC: &mut MultilinearPolynomial<E::Scalar>,
poly_z: &mut MultilinearPolynomial<E::Scalar>,
verifier_circuit: &mut SpartanVerifierCircuit<E>,
state: &mut MultiRoundState<E>,
vc_shape: &SplitMultiRoundR1CSShape<E>,
vc_ck: &CommitmentKey<E>,
transcript: &mut E::TE,
start_round: usize,
) -> Result<(Vec<E::Scalar>, Vec<E::Scalar>), SpartanError> {
let mut r_y: Vec<E::Scalar> = Vec::with_capacity(num_rounds);
let mut claim_current_round = *claim;
for j in 0..num_rounds {
let (eval0, eval2) = Self::compute_eval_points_quad(poly_ABC, poly_z);
let evals = vec![eval0, claim_current_round - eval0, eval2];
let poly = UniPoly::from_evals(&evals)?;
verifier_circuit.inner_polys[j] = [poly.coeffs[0], poly.coeffs[1], poly.coeffs[2]];
let chals = SatisfyingAssignment::<E>::process_round(
state,
vc_shape,
vc_ck,
verifier_circuit,
start_round + j,
transcript,
)?;
r_y.push(chals[0]);
rayon::join(
|| poly_ABC.bind_poly_var_top(&chals[0]),
|| poly_z.bind_poly_var_top(&chals[0]),
);
claim_current_round = poly.evaluate(&chals[0]);
}
Ok((r_y, vec![poly_ABC[0], poly_z[0]]))
}
pub fn prove_quad_batched_zk(
claims: &[E::Scalar; 2],
num_rounds: usize,
poly_A_0: &mut MultilinearPolynomial<E::Scalar>,
poly_A_1: &mut MultilinearPolynomial<E::Scalar>,
poly_B_0: &mut MultilinearPolynomial<E::Scalar>,
poly_B_1: &mut MultilinearPolynomial<E::Scalar>,
verifier_circuit: &mut NeutronNovaVerifierCircuit<E>,
state: &mut MultiRoundState<E>,
vc_shape: &SplitMultiRoundR1CSShape<E>,
vc_ck: &CommitmentKey<E>,
transcript: &mut E::TE,
start_round: usize,
) -> Result<(Vec<E::Scalar>, Vec<E::Scalar>), SpartanError> {
let mut r_y: Vec<E::Scalar> = Vec::with_capacity(num_rounds);
let mut claim_step_round = claims[0];
let mut claim_core_round = claims[1];
for j in 0..num_rounds {
let ((eval0_s, eval2_s), (eval0_c, eval2_c)) = rayon::join(
|| Self::compute_eval_points_quad(poly_A_0, poly_B_0),
|| Self::compute_eval_points_quad(poly_A_1, poly_B_1),
);
let evals_s = vec![eval0_s, claim_step_round - eval0_s, eval2_s];
let poly_s = UniPoly::from_evals(&evals_s)?;
let coeffs_step = [poly_s.coeffs[0], poly_s.coeffs[1], poly_s.coeffs[2]];
let evals_c = vec![eval0_c, claim_core_round - eval0_c, eval2_c];
let poly_c = UniPoly::from_evals(&evals_c)?;
let coeffs_core = [poly_c.coeffs[0], poly_c.coeffs[1], poly_c.coeffs[2]];
verifier_circuit.inner_polys_step[j] = coeffs_step;
verifier_circuit.inner_polys_core[j] = coeffs_core;
let chals = SatisfyingAssignment::<E>::process_round(
state,
vc_shape,
vc_ck,
verifier_circuit,
start_round + j,
transcript,
)?;
let r_j = chals[0];
r_y.push(r_j);
rayon::join(
|| {
rayon::join(
|| poly_A_0.bind_poly_var_top(&r_j),
|| poly_B_0.bind_poly_var_top(&r_j),
);
},
|| {
rayon::join(
|| poly_A_1.bind_poly_var_top(&r_j),
|| poly_B_1.bind_poly_var_top(&r_j),
);
},
);
claim_step_round = poly_s.evaluate(&r_j);
claim_core_round = poly_c.evaluate(&r_j);
}
Ok((
r_y,
vec![poly_A_0[0], poly_A_1[0], poly_B_0[0], poly_B_1[0]],
))
}
pub fn prove_cubic_with_additive_term_batched_zk(
num_rounds: usize,
pow_tau_left: &mut MultilinearPolynomial<E::Scalar>,
pow_tau_right: &MultilinearPolynomial<E::Scalar>,
poly_A_step: &mut MultilinearPolynomial<E::Scalar>,
poly_A_core: &mut MultilinearPolynomial<E::Scalar>,
poly_B_step: &mut MultilinearPolynomial<E::Scalar>,
poly_B_core: &mut MultilinearPolynomial<E::Scalar>,
poly_C_step: &mut MultilinearPolynomial<E::Scalar>,
poly_C_core: &mut MultilinearPolynomial<E::Scalar>,
verifier_circuit: &mut NeutronNovaVerifierCircuit<E>,
state: &mut MultiRoundState<E>,
vc_shape: &SplitMultiRoundR1CSShape<E>,
vc_ck: &CommitmentKey<E>,
transcript: &mut E::TE,
start_round: usize,
) -> Result<Vec<E::Scalar>, SpartanError> {
let mut base_tau = E::Scalar::ONE;
let mut len_pow_tau = pow_tau_left.Z.len() * pow_tau_right.Z.len();
let mut r_x: Vec<E::Scalar> = Vec::with_capacity(num_rounds);
let mut claim_step = verifier_circuit.t_out_step;
let mut claim_core = E::Scalar::ZERO;
for i in 0..num_rounds {
let ((mut eval0_s, mut eval2_s, mut eval3_s), (mut eval0_c, mut eval2_c, mut eval3_c)) =
rayon::join(
|| {
Self::compute_eval_points_cubic_with_additive_term_with_outer_pow(
pow_tau_left,
pow_tau_right,
poly_A_step,
poly_B_step,
poly_C_step,
)
},
|| {
Self::compute_eval_points_cubic_with_additive_term_with_outer_pow(
pow_tau_left,
pow_tau_right,
poly_A_core,
poly_B_core,
poly_C_core,
)
},
);
eval0_s *= base_tau;
eval2_s *= base_tau;
eval3_s *= base_tau;
eval0_c *= base_tau;
eval2_c *= base_tau;
eval3_c *= base_tau;
let evals_s = vec![eval0_s, claim_step - eval0_s, eval2_s, eval3_s];
let poly_s = UniPoly::from_evals(&evals_s)?;
let coeffs_step = [
poly_s.coeffs[0],
poly_s.coeffs[1],
poly_s.coeffs[2],
poly_s.coeffs[3],
];
let evals_c = vec![eval0_c, claim_core - eval0_c, eval2_c, eval3_c];
let poly_c = UniPoly::from_evals(&evals_c)?;
let coeffs_core = [
poly_c.coeffs[0],
poly_c.coeffs[1],
poly_c.coeffs[2],
poly_c.coeffs[3],
];
verifier_circuit.outer_polys_step[i] = coeffs_step;
verifier_circuit.outer_polys_core[i] = coeffs_core;
let chals = SatisfyingAssignment::<E>::process_round(
state,
vc_shape,
vc_ck,
verifier_circuit,
start_round + i,
transcript,
)?;
let r_i = chals[0];
r_x.push(r_i);
claim_step = poly_s.evaluate(&r_i);
claim_core = poly_c.evaluate(&r_i);
rayon::join(
|| {
rayon::join(
|| poly_A_step.bind_poly_var_top(&r_i),
|| poly_A_core.bind_poly_var_top(&r_i),
);
},
|| {
rayon::join(
|| {
rayon::join(
|| poly_B_step.bind_poly_var_top(&r_i),
|| poly_B_core.bind_poly_var_top(&r_i),
);
},
|| {
rayon::join(
|| poly_C_step.bind_poly_var_top(&r_i),
|| poly_C_core.bind_poly_var_top(&r_i),
);
},
);
},
);
len_pow_tau >>= 1;
let one = E::Scalar::ONE;
let left = pow_tau_left.Z.len();
let pow = pow_tau_left.Z[len_pow_tau % left] * pow_tau_right.Z[len_pow_tau / left];
base_tau *= (pow - one) * r_i + one;
}
pow_tau_left.Z[0] = base_tau;
Ok(r_x)
}
}
pub(crate) mod eq_sumcheck {
use crate::{
big_num::DelayedReduction, polys::multilinear::MultilinearPolynomial, traits::Engine,
};
use ff::{Field, PrimeField};
use num_traits::Zero;
use rayon::prelude::*;
pub struct EqSumCheckInstance<E: Engine> {
init_num_vars: usize,
first_half: usize,
second_half: usize,
round: usize,
taus: Vec<E::Scalar>,
eval_eq_left: E::Scalar,
poly_eq_left: Vec<Vec<E::Scalar>>,
poly_eq_right: Vec<Vec<E::Scalar>>,
eq_tau_0_2_3: Vec<(E::Scalar, E::Scalar, E::Scalar)>,
}
impl<E: Engine> EqSumCheckInstance<E> {
pub fn new(taus: Vec<E::Scalar>) -> Self {
let l = taus.len();
let first_half = l / 2;
let compute_eq_polynomials = |taus: Vec<&E::Scalar>| -> Vec<Vec<E::Scalar>> {
let len = taus.len();
let mut result = Vec::with_capacity(len + 1);
result.push(vec![E::Scalar::ONE]);
for i in 0..len {
let tau = taus[i];
let prev = &result[i];
let mut v_next = prev.to_vec();
v_next.par_extend(prev.par_iter().map(|v| *v * tau));
let (first, last) = v_next.split_at_mut(prev.len());
first.par_iter_mut().zip(last).for_each(|(a, b)| *a -= *b);
result.push(v_next);
}
result
};
let (left_taus, right_taus) = taus.split_at(first_half);
let left_taus = left_taus.iter().skip(1).rev().collect::<Vec<_>>();
let right_taus = right_taus.iter().rev().collect::<Vec<_>>();
let (poly_eq_left, poly_eq_right) = rayon::join(
|| compute_eq_polynomials(left_taus),
|| compute_eq_polynomials(right_taus),
);
let f2 = E::Scalar::ONE.double();
let f1 = E::Scalar::ONE;
let eq_tau_0_2_3 = taus
.par_iter()
.map(|tau| {
let tau2 = tau.double();
let tau3 = tau2 + tau;
let tau5 = tau3 + tau2;
(f1 - tau, tau3 - f1, tau5 - f2)
})
.collect::<Vec<_>>();
Self {
init_num_vars: l,
first_half,
second_half: l - first_half,
round: 1, taus,
eval_eq_left: E::Scalar::ONE,
poly_eq_left,
poly_eq_right,
eq_tau_0_2_3,
}
}
#[inline]
pub fn evaluation_points_cubic_with_three_inputs(
&self,
round_idx: usize,
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
poly_C: &MultilinearPolynomial<E::Scalar>,
) -> (E::Scalar, E::Scalar, E::Scalar) {
debug_assert_eq!(poly_A.Z.len() % 2, 0);
type Acc<S> = <S as DelayedReduction<S>>::Accumulator;
let in_first_half = self.round < self.first_half;
let half_p = poly_A.Z.len() / 2;
let (mut eval_0, mut eval_2, mut eval_3) = if in_first_half {
let (poly_eq_left, poly_eq_right, second_half) = self.poly_eqs_first_half();
let eq_out_len = poly_eq_left.len();
let (acc_0, acc_2, acc_3) = (0..eq_out_len)
.into_par_iter()
.fold(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut outer_acc, x_out| {
let e_out = &poly_eq_left[x_out];
let mut inner_0 = Acc::<E::Scalar>::zero();
let mut inner_2 = Acc::<E::Scalar>::zero();
let mut inner_3 = Acc::<E::Scalar>::zero();
for (x_in, e_in) in poly_eq_right.iter().enumerate() {
let id = (x_out << second_half) | x_in;
let (zero_a, one_a) = (&poly_A.Z[id], &poly_A.Z[id + half_p]);
let (zero_b, one_b) = (&poly_B.Z[id], &poly_B.Z[id + half_p]);
let (zero_c, one_c) = (&poly_C.Z[id], &poly_C.Z[id + half_p]);
let (q0, q2, q3) = eval_one_case_cubic_three_inputs(
round_idx, zero_a, one_a, zero_b, one_b, zero_c, one_c,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_0,
e_in,
&q0,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_2,
e_in,
&q2,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut inner_3,
e_in,
&q3,
);
}
let inner_0_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_0);
let inner_2_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_2);
let inner_3_red = <E::Scalar as DelayedReduction<E::Scalar>>::reduce(&inner_3);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.0,
e_out,
&inner_0_red,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.1,
e_out,
&inner_2_red,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut outer_acc.2,
e_out,
&inner_3_red,
);
outer_acc
},
)
.reduce(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut a, b| {
a.0 += b.0;
a.1 += b.1;
a.2 += b.2;
a
},
);
(
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_0),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_2),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_3),
)
} else {
let poly_eq_right = self.poly_eq_right_last_half();
let (acc_0, acc_2, acc_3) = (0..half_p)
.into_par_iter()
.fold(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut acc, id| {
let e = &poly_eq_right[id];
let (zero_a, one_a) = (&poly_A.Z[id], &poly_A.Z[id + half_p]);
let (zero_b, one_b) = (&poly_B.Z[id], &poly_B.Z[id + half_p]);
let (zero_c, one_c) = (&poly_C.Z[id], &poly_C.Z[id + half_p]);
let (q0, q2, q3) = eval_one_case_cubic_three_inputs(
round_idx, zero_a, one_a, zero_b, one_b, zero_c, one_c,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.0, e, &q0,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.1, e, &q2,
);
<E::Scalar as DelayedReduction<E::Scalar>>::unreduced_multiply_accumulate(
&mut acc.2, e, &q3,
);
acc
},
)
.reduce(
|| {
(
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
Acc::<E::Scalar>::zero(),
)
},
|mut a, b| {
a.0 += b.0;
a.1 += b.1;
a.2 += b.2;
a
},
);
(
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_0),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_2),
<E::Scalar as DelayedReduction<E::Scalar>>::reduce(&acc_3),
)
};
self.update_evals(&mut eval_0, &mut eval_2, &mut eval_3);
(eval_0, eval_2, eval_3)
}
#[inline]
pub fn bound(&mut self, r: &E::Scalar) {
let tau = self.taus[self.round - 1];
self.eval_eq_left *= E::Scalar::ONE - tau - r + (*r * tau).double();
self.round += 1;
}
#[inline]
fn update_evals(&self, eval_0: &mut E::Scalar, eval_2: &mut E::Scalar, eval_3: &mut E::Scalar) {
let p = self.eval_eq_left;
let eq_tau_0_2_3 = self.eq_tau_0_2_3[self.round - 1];
let eq_tau_0_p = eq_tau_0_2_3.0 * p;
let eq_tau_2_p = eq_tau_0_2_3.1 * p;
let eq_tau_3_p = eq_tau_0_2_3.2 * p;
*eval_0 *= eq_tau_0_p;
*eval_2 *= eq_tau_2_p;
*eval_3 *= eq_tau_3_p;
}
#[inline]
fn poly_eqs_first_half(&self) -> (&Vec<E::Scalar>, &Vec<E::Scalar>, usize) {
let second_half = self.second_half;
let poly_eq_left = &self.poly_eq_left[self.first_half - self.round];
let poly_eq_right = &self.poly_eq_right[second_half];
debug_assert_eq!(poly_eq_right.len(), 1 << second_half);
(poly_eq_left, poly_eq_right, second_half)
}
#[inline]
fn poly_eq_right_last_half(&self) -> &Vec<E::Scalar> {
&self.poly_eq_right[self.init_num_vars - self.round]
}
}
#[inline]
fn eval_one_case_cubic_three_inputs<Scalar: PrimeField>(
round_idx: usize,
zero_a: &Scalar,
one_a: &Scalar,
zero_b: &Scalar,
one_b: &Scalar,
zero_c: &Scalar,
one_c: &Scalar,
) -> (Scalar, Scalar, Scalar) {
let eval_0 = if round_idx == 0 {
Scalar::ZERO
} else {
*zero_a * *zero_b - *zero_c
};
let double_one_a = one_a.double();
let double_one_b = one_b.double();
let double_one_c = one_c.double();
let eval_2 = {
let point_a = double_one_a - *zero_a;
let point_b = double_one_b - *zero_b;
let point_c = double_one_c - *zero_c;
point_a * point_b - point_c
};
let eval_3 = {
let point_a = double_one_a + one_a - zero_a.double();
let point_b = double_one_b + one_b - zero_b.double();
let point_c = double_one_c + one_c - zero_c.double();
point_a * point_b - point_c
};
(eval_0, eval_2, eval_3)
}
}
#[cfg(test)]
mod perf_tests {
use super::*;
use crate::{
big_num::DelayedReduction, polys::multilinear::MultilinearPolynomial, start_span,
traits::Engine,
};
use ff::Field;
use rand::{SeedableRng, rngs::StdRng};
use tracing::info;
use tracing_subscriber::EnvFilter;
#[cfg(debug_assertions)]
const TEST_SIZES: &[usize] = &[16, 18];
#[cfg(not(debug_assertions))]
const TEST_SIZES: &[usize] = &[16, 18, 20, 22, 24];
fn test_first_round_spartan_sumcheck_with<E: Engine>()
where
E::Scalar: DelayedReduction<E::Scalar>,
{
const SEED: u64 = 0xDEADBEEF;
let field_name = std::any::type_name::<E::Scalar>()
.split("::")
.last()
.unwrap_or("unknown");
for &num_vars in TEST_SIZES {
let len = 1 << num_vars;
let mut rng = StdRng::seed_from_u64(SEED);
let az: Vec<E::Scalar> = (0..len).map(|_| E::Scalar::random(&mut rng)).collect();
let bz: Vec<E::Scalar> = (0..len).map(|_| E::Scalar::random(&mut rng)).collect();
let cz: Vec<E::Scalar> = (0..len).map(|_| E::Scalar::random(&mut rng)).collect();
let taus: Vec<E::Scalar> = (0..num_vars).map(|_| E::Scalar::random(&mut rng)).collect();
let mut poly_az = MultilinearPolynomial::new(az);
let mut poly_bz = MultilinearPolynomial::new(bz);
let mut poly_cz = MultilinearPolynomial::new(cz);
let mut transcript = E::TE::new(b"perf_test");
let (_span, t) = start_span!("sumcheck_prove", field = field_name, num_vars = num_vars);
let (proof, _r, _evals) = SumcheckProof::<E>::prove_cubic_with_three_inputs(
&E::Scalar::ZERO,
taus,
&mut poly_az,
&mut poly_bz,
&mut poly_cz,
&mut transcript,
)
.expect("proof generation should succeed");
info!(field = field_name, num_vars, n = len, ms = ?t.elapsed().as_millis(), "completed");
let mut verifier_transcript = E::TE::new(b"perf_test");
proof
.verify(E::Scalar::ZERO, num_vars, 3, &mut verifier_transcript)
.expect("proof verification should succeed");
}
}
#[test]
fn test_first_round_spartan_sumcheck() {
let _ = tracing_subscriber::fmt()
.with_target(false)
.with_ansi(true)
.with_env_filter(EnvFilter::from_default_env())
.try_init();
use crate::provider::Bn254Engine;
test_first_round_spartan_sumcheck_with::<Bn254Engine>();
#[cfg(not(debug_assertions))]
{
use crate::provider::{PallasHyraxEngine, T256HyraxEngine};
test_first_round_spartan_sumcheck_with::<PallasHyraxEngine>();
test_first_round_spartan_sumcheck_with::<T256HyraxEngine>();
}
}
fn test_inner_sumcheck_with<E: Engine>()
where
E::Scalar: DelayedReduction<E::Scalar>,
{
const SEED: u64 = 0xDEADBEEF;
let field_name = std::any::type_name::<E::Scalar>()
.split("::")
.last()
.unwrap_or("unknown");
for &num_vars in TEST_SIZES {
let len = 1 << num_vars;
let mut rng = StdRng::seed_from_u64(SEED);
let poly_a: Vec<E::Scalar> = (0..len).map(|_| E::Scalar::random(&mut rng)).collect();
let poly_b: Vec<E::Scalar> = (0..len).map(|_| E::Scalar::random(&mut rng)).collect();
let claim: E::Scalar = poly_a.par_iter().zip(&poly_b).map(|(a, b)| *a * *b).sum();
let mut poly_a = MultilinearPolynomial::new(poly_a);
let mut poly_b = MultilinearPolynomial::new(poly_b);
let mut transcript = E::TE::new(b"test_inner_sumcheck");
let (_span, t) = start_span!("prove_quad", field = field_name, num_vars = num_vars);
let (proof, _r, _evals) =
SumcheckProof::<E>::prove_quad(&claim, num_vars, &mut poly_a, &mut poly_b, &mut transcript)
.expect("proof generation should succeed");
info!(field = field_name, num_vars, n = len, ms = ?t.elapsed().as_millis(), "completed");
let mut verifier_transcript = E::TE::new(b"test_inner_sumcheck");
proof
.verify(claim, num_vars, 2, &mut verifier_transcript)
.expect("proof verification should succeed");
}
}
#[test]
fn test_inner_sumcheck() {
let _ = tracing_subscriber::fmt()
.with_target(false)
.with_ansi(true)
.with_env_filter(EnvFilter::from_default_env())
.try_init();
use crate::provider::Bn254Engine;
test_inner_sumcheck_with::<Bn254Engine>();
#[cfg(not(debug_assertions))]
{
use crate::provider::{PallasHyraxEngine, T256HyraxEngine};
test_inner_sumcheck_with::<PallasHyraxEngine>();
test_inner_sumcheck_with::<T256HyraxEngine>();
}
}
}