use crate::{
errors::SpartanError,
polys::{
multilinear::MultilinearPolynomial,
univariate::{CompressedUniPoly, UniPoly},
},
start_span,
traits::{Engine, transcript::TranscriptEngineTrait},
};
use ff::Field;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use std::time::Instant;
use tracing::{info, info_span};
const PAR_THRESHOLD: usize = 4 << 10;
pub fn par_for<R, Map, Red, Id>(len: usize, map: Map, reduce: Red, identity: Id) -> R
where
R: Send, Map: Fn(usize) -> R + Sync + Send,
Red: Fn(R, R) -> R + Sync + Send,
Id: Fn() -> R + Sync + Send,
{
if len == 0 {
return identity();
}
let in_rayon_ctx = rayon::current_thread_index().is_some();
if len < PAR_THRESHOLD || in_rayon_ctx {
let mut acc = identity();
for i in 0..len {
let v = map(i);
acc = reduce(acc, v);
}
acc
} else {
(0..len).into_par_iter().map(map).reduce(identity, reduce)
}
}
#[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<F>(
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
comb_func: &F,
) -> (E::Scalar, E::Scalar)
where
F: Fn(&E::Scalar, &E::Scalar) -> E::Scalar + Sync,
{
let len = poly_A.Z.len() / 2;
par_for(
len,
|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];
let eval0 = comb_func(&a_low, &b_low);
let a_bound = a_high + a_high - a_low;
let b_bound = b_high + b_high - b_low;
let eval2 = comb_func(&a_bound, &b_bound);
(eval0, eval2)
},
|mut acc, val| {
acc.0 += val.0;
acc.1 += val.1;
acc
},
|| (E::Scalar::ZERO, E::Scalar::ZERO),
)
}
pub fn prove_quad<F>(
claim: &E::Scalar,
num_rounds: usize,
poly_A: &mut MultilinearPolynomial<E::Scalar>,
poly_B: &mut MultilinearPolynomial<E::Scalar>,
comb_func: F,
transcript: &mut E::TE,
) -> Result<(Self, Vec<E::Scalar>, Vec<E::Scalar>), SpartanError>
where
F: Fn(&E::Scalar, &E::Scalar) -> E::Scalar + Sync,
{
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, &comb_func);
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]
pub fn compute_eval_points_cubic_with_additive_term<F>(
poly_A: &MultilinearPolynomial<E::Scalar>,
poly_B: &MultilinearPolynomial<E::Scalar>,
poly_C: &MultilinearPolynomial<E::Scalar>,
poly_D: &MultilinearPolynomial<E::Scalar>,
comb_func: &F,
) -> (E::Scalar, E::Scalar, E::Scalar)
where
F: Fn(&E::Scalar, &E::Scalar, &E::Scalar, &E::Scalar) -> E::Scalar + Sync,
{
let len = poly_A.Z.len() / 2;
par_for(
len,
|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 eval_point_0 = comb_func(&a_low, &b_low, &c_low, &d_low);
let poly_A_bound_point = a_high + a_high - a_low;
let poly_B_bound_point = b_high + b_high - b_low;
let poly_C_bound_point = c_high + c_high - c_low;
let poly_D_bound_point = d_high + d_high - d_low;
let eval_point_2 = comb_func(
&poly_A_bound_point,
&poly_B_bound_point,
&poly_C_bound_point,
&poly_D_bound_point,
);
let poly_A_bound_point = poly_A_bound_point + a_high - a_low;
let poly_B_bound_point = poly_B_bound_point + b_high - b_low;
let poly_C_bound_point = poly_C_bound_point + c_high - c_low;
let poly_D_bound_point = poly_D_bound_point + d_high - d_low;
let eval_point_3 = comb_func(
&poly_A_bound_point,
&poly_B_bound_point,
&poly_C_bound_point,
&poly_D_bound_point,
);
(eval_point_0, eval_point_2, eval_point_3)
},
|mut acc, val| {
acc.0 += val.0;
acc.1 += val.1;
acc.2 += val.2;
acc
},
|| (E::Scalar::ZERO, E::Scalar::ZERO, E::Scalar::ZERO),
)
}
#[allow(clippy::too_many_arguments)]
pub fn prove_cubic_with_additive_term<F>(
claim: &E::Scalar,
num_rounds: usize,
poly_A: &mut MultilinearPolynomial<E::Scalar>,
poly_B: &mut MultilinearPolynomial<E::Scalar>,
poly_C: &mut MultilinearPolynomial<E::Scalar>,
poly_D: &mut MultilinearPolynomial<E::Scalar>,
comb_func: F,
transcript: &mut E::TE,
) -> Result<(Self, Vec<E::Scalar>, Vec<E::Scalar>), SpartanError>
where
F: Fn(&E::Scalar, &E::Scalar, &E::Scalar, &E::Scalar) -> E::Scalar + Sync,
{
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_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) =
Self::compute_eval_points_cubic_with_additive_term(
poly_A, poly_B, poly_C, poly_D, &comb_func,
);
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(
|| {
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),
|| poly_D.bind_poly_var_top(&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], poly_D[0]],
))
}
}