use std::ops::Add;
use p3_field::{Field, PrimeCharacteristicRing};
use p3_util::log2_strict_usize;
use tracing::{debug, instrument};
use crate::{
proof::GkrLayerClaims,
prover::{
error::LogupZerocheckError,
poly::evals_eq_hypercube,
sumcheck::{fold_mle_evals, sumcheck_round_poly_evals},
ColMajorMatrix,
},
FiatShamirTranscript, StarkProtocolConfig,
};
pub struct FracSumcheckProof<SC: StarkProtocolConfig> {
pub fractional_sum: (SC::EF, SC::EF),
pub claims_per_layer: Vec<GkrLayerClaims<SC>>,
pub sumcheck_polys: Vec<Vec<[SC::EF; 3]>>,
}
#[derive(Clone, Copy, Debug, Default, derive_new::new)]
#[repr(C)]
pub struct Frac<EF> {
pub p: EF,
pub q: EF,
}
impl<EF: Field> Add<Frac<EF>> for Frac<EF> {
type Output = Frac<EF>;
fn add(self, other: Frac<EF>) -> Self::Output {
Frac {
p: self.p * other.q + self.q * other.p,
q: self.q * other.q,
}
}
}
#[instrument(level = "info", skip_all)]
pub fn fractional_sumcheck<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
transcript: &mut TS,
evals: &[Frac<SC::EF>],
assert_zero: bool,
) -> Result<(FracSumcheckProof<SC>, Vec<SC::EF>), LogupZerocheckError> {
if evals.is_empty() {
return Ok((
FracSumcheckProof {
fractional_sum: (SC::EF::ZERO, SC::EF::ONE),
claims_per_layer: vec![],
sumcheck_polys: vec![],
},
vec![],
));
}
let total_rounds = log2_strict_usize(evals.len());
let mut sumcheck_polys = Vec::with_capacity(total_rounds);
let mut tree_evals: Vec<Frac<SC::EF>> = vec![Frac::default(); 2 << total_rounds];
tree_evals[(1 << total_rounds)..].copy_from_slice(evals);
for node_idx in (1..(1 << total_rounds)).rev() {
tree_evals[node_idx] = tree_evals[2 * node_idx] + tree_evals[2 * node_idx + 1];
}
let frac_sum = tree_evals[1];
if assert_zero {
if frac_sum.p != SC::EF::ZERO {
return Err(LogupZerocheckError::NonZeroRootSum);
}
} else {
transcript.observe_ext(frac_sum.p);
}
transcript.observe_ext(frac_sum.q);
let mut claims_per_layer: Vec<GkrLayerClaims<SC>> = Vec::with_capacity(total_rounds);
claims_per_layer.push(GkrLayerClaims::<SC> {
p_xi_0: tree_evals[2].p,
q_xi_0: tree_evals[2].q,
p_xi_1: tree_evals[3].p,
q_xi_1: tree_evals[3].q,
});
transcript.observe_ext(claims_per_layer[0].p_xi_0);
transcript.observe_ext(claims_per_layer[0].q_xi_0);
transcript.observe_ext(claims_per_layer[0].p_xi_1);
transcript.observe_ext(claims_per_layer[0].q_xi_1);
let mu_1 = transcript.sample_ext();
debug!(gkr_round = 0, mu = %mu_1);
let mut xi_prev = vec![mu_1];
for round in 1..total_rounds {
let eval_size = 1 << round;
let lambda = transcript.sample_ext();
debug!(gkr_round = round, %lambda);
let mut pq_j_evals = SC::EF::zero_vec(4 * eval_size);
let segment = &tree_evals[2 * eval_size..4 * eval_size];
for x in 0..eval_size {
pq_j_evals[x] = segment[2 * x].p;
pq_j_evals[eval_size + x] = segment[2 * x].q;
pq_j_evals[2 * eval_size + x] = segment[2 * x + 1].p;
pq_j_evals[3 * eval_size + x] = segment[2 * x + 1].q;
}
let mut pq_j_evals = ColMajorMatrix::new(pq_j_evals, 4);
let mut eq_xis = ColMajorMatrix::new(evals_eq_hypercube(&xi_prev), 1);
let (round_polys_eval, rho) = {
let n = round;
let mut round_polys_eval = Vec::with_capacity(n);
let mut r_vec = Vec::with_capacity(n);
for sumcheck_round in 0..n {
let [s_evals] = sumcheck_round_poly_evals(
n - sumcheck_round,
3,
&[eq_xis.as_view(), pq_j_evals.as_view()],
|_x, _y, row| {
let eq_xi = row[0][0];
let &[p_j0, q_j0, p_j1, q_j1] = row[1].as_slice() else {
unreachable!("pq_j_evals always has 4 columns")
};
let p_prev = p_j0 * q_j1 + p_j1 * q_j0;
let q_prev = q_j0 * q_j1;
[eq_xi * (p_prev + lambda * q_prev)]
},
);
let s_evals: [SC::EF; 3] = s_evals.try_into().unwrap();
for &eval in &s_evals {
transcript.observe_ext(eval);
}
round_polys_eval.push(s_evals);
let r_round = transcript.sample_ext();
pq_j_evals = fold_mle_evals(pq_j_evals, r_round);
eq_xis = fold_mle_evals(eq_xis, r_round);
r_vec.push(r_round);
debug!(gkr_round = round, %sumcheck_round, %r_round);
}
(round_polys_eval, r_vec)
};
claims_per_layer.push(GkrLayerClaims::<SC> {
p_xi_0: pq_j_evals.column(0)[0],
q_xi_0: pq_j_evals.column(1)[0],
p_xi_1: pq_j_evals.column(2)[0],
q_xi_1: pq_j_evals.column(3)[0],
});
transcript.observe_ext(claims_per_layer[round].p_xi_0);
transcript.observe_ext(claims_per_layer[round].q_xi_0);
transcript.observe_ext(claims_per_layer[round].p_xi_1);
transcript.observe_ext(claims_per_layer[round].q_xi_1);
let mu = transcript.sample_ext();
debug!(gkr_round = round, %mu);
xi_prev = [vec![mu], rho].concat();
sumcheck_polys.push(round_polys_eval);
}
Ok((
FracSumcheckProof {
fractional_sum: (frac_sum.p, frac_sum.q),
claims_per_layer,
sumcheck_polys,
},
xi_prev,
))
}