use std::{cmp::max, iter::zip};
use itertools::Itertools;
use p3_dft::TwoAdicSubgroupDft;
use p3_field::{ExtensionField, Field, PrimeCharacteristicRing, TwoAdicField};
use p3_matrix::dense::RowMajorMatrix;
use p3_maybe_rayon::prelude::*;
use p3_util::log2_strict_usize;
use tracing::{debug, info_span, instrument};
use crate::{
calculate_n_logup,
dft::Radix2BowersSerial,
poly_common::{eq_uni_poly, UnivariatePoly},
proof::{column_openings_by_rot, BatchConstraintProof, GkrProof},
prover::{
error::LogupZerocheckError,
fractional_sumcheck_gkr::{fractional_sumcheck, Frac},
poly::eq_sharp_uni_poly,
stacked_pcs::StackedLayout,
sumcheck::sumcheck_round0_deg,
ColMajorMatrix, CpuColMajorBackend, DeviceMultiStarkProvingKey, MatrixDimensions,
MatrixView, ProverBackend, ProvingContext,
},
FiatShamirTranscript, StarkProtocolConfig,
};
mod cpu;
mod evaluator;
pub mod fractional_sumcheck_gkr;
mod single;
pub use cpu::LogupZerocheckCpu;
pub use single::*;
#[instrument(level = "info", skip_all)]
#[allow(clippy::type_complexity)]
pub fn prove_zerocheck_and_logup<SC: StarkProtocolConfig, TS>(
transcript: &mut TS,
mpk: &DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
ctx: &ProvingContext<CpuColMajorBackend<SC>>,
) -> Result<(GkrProof<SC>, BatchConstraintProof<SC>, Vec<SC::EF>), LogupZerocheckError>
where
TS: FiatShamirTranscript<SC>,
SC::F: TwoAdicField,
SC::EF: TwoAdicField + ExtensionField<SC::F>,
CpuColMajorBackend<SC>: ProverBackend<Val = SC::F, Matrix = ColMajorMatrix<SC::F>>,
{
let l_skip = mpk.params.l_skip;
let constraint_degree = mpk.max_constraint_degree;
let num_traces = ctx.per_trace.len();
let n_max = log2_strict_usize(ctx.per_trace[0].1.common_main.height()).saturating_sub(l_skip);
let mut total_interactions = 0u64;
let interactions_meta: Vec<_> = ctx
.per_trace
.iter()
.map(|(air_idx, trace_ctx)| {
let pk = &mpk.per_air[*air_idx];
let num_interactions = pk.vk.symbolic_constraints.interactions.len();
let height = trace_ctx.common_main.height();
let log_height = log2_strict_usize(height);
let log_lifted_height = log_height.max(l_skip);
total_interactions += (num_interactions as u64) << log_lifted_height;
(num_interactions, log_lifted_height)
})
.collect();
let n_logup = calculate_n_logup(l_skip, total_interactions);
debug!(%n_logup);
let interactions_layout = StackedLayout::new(0, l_skip + n_logup, interactions_meta)?;
let logup_pow_witness = transcript.grind(mpk.params.logup.pow_bits);
let alpha_logup = transcript.sample_ext();
let beta_logup = transcript.sample_ext();
debug!(%alpha_logup, %beta_logup);
let mut prover = LogupZerocheckCpu::new(
mpk,
ctx,
n_logup,
interactions_layout,
alpha_logup,
beta_logup,
)?;
let has_interactions = !prover.interactions_layout.sorted_cols.is_empty();
let gkr_input_evals = if !has_interactions {
vec![]
} else {
let unstacked_interaction_evals = prover
.eval_helpers
.par_iter()
.enumerate()
.map(|(trace_idx, helper)| {
let trace_ctx = &ctx.per_trace[trace_idx].1;
let mats = helper.view_mats(trace_ctx);
let height = trace_ctx.common_main.height();
(0..height)
.into_par_iter()
.map(|i| {
let mut row_parts = Vec::with_capacity(mats.len() + 1);
let is_first = SC::F::from_bool(i == 0);
let is_transition = SC::F::from_bool(i != height - 1);
let is_last = SC::F::from_bool(i == height - 1);
let sels = vec![is_first, is_transition, is_last];
row_parts.push(sels);
for (mat, is_rot) in &mats {
let offset = usize::from(*is_rot);
row_parts.push(
(0..mat.width())
.map(|j| unsafe {
*mat.get_unchecked((i + offset) % height, j)
})
.collect_vec(),
);
}
helper.eval_interactions(&row_parts, &prover.beta_pows)
})
.collect::<Vec<_>>()
})
.collect::<Vec<_>>();
let mut evals = vec![Frac::default(); 1 << (l_skip + n_logup)];
for (trace_idx, interaction_idx, s) in
prover.interactions_layout.sorted_cols.iter().copied()
{
let pq_evals = &unstacked_interaction_evals[trace_idx];
let height = pq_evals.len();
debug_assert_eq!(s.col_idx, 0);
debug_assert_eq!(1 << s.log_height(), s.len(0));
debug_assert_eq!(s.len(0) % height, 0);
let norm_factor_denom = s.len(0) / height;
let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
evals[s.row_idx..s.row_idx + s.len(0)]
.chunks_exact_mut(height)
.for_each(|evals| {
evals
.par_iter_mut()
.zip(pq_evals)
.for_each(|(pq_eval, evals_at_z)| {
let (mut numer, denom) = evals_at_z[interaction_idx];
numer *= norm_factor;
*pq_eval = Frac::new(numer.into(), denom);
});
});
}
evals.par_iter_mut().for_each(|frac| frac.q += alpha_logup);
evals
};
let (frac_sum_proof, mut xi) =
fractional_sumcheck::<SC, _>(transcript, &gkr_input_evals, true)?;
let n_global = max(n_max, n_logup);
debug!(%n_global);
while xi.len() != l_skip + n_global {
xi.push(transcript.sample_ext());
}
debug!(?xi);
prover.xi = xi;
let mut sumcheck_round_polys = Vec::with_capacity(n_max);
let mut r = Vec::with_capacity(n_max + 1);
let lambda = transcript.sample_ext();
debug!(%lambda);
let sp_0_polys = prover.sumcheck_uni_round0_polys(ctx, lambda)?;
let sp_0_deg = sumcheck_round0_deg(l_skip, constraint_degree);
let s_deg = constraint_degree + 1;
let s_0_deg = sumcheck_round0_deg(l_skip, s_deg);
let large_uni_domain = (s_0_deg + 1).next_power_of_two();
let dft = Radix2BowersSerial;
let s_0_logup_polys = {
let eq_sharp_uni = eq_sharp_uni_poly(&prover.xi[..l_skip]);
let mut eq_coeffs = eq_sharp_uni.into_coeffs();
eq_coeffs.resize(large_uni_domain, SC::EF::ZERO);
let eq_evals = dft.dft(eq_coeffs);
let width = 2 * num_traces;
let mut sp_coeffs_mat = SC::EF::zero_vec(width * large_uni_domain);
for (i, coeffs) in sp_0_polys[..2 * num_traces].iter().enumerate() {
for (j, &c_j) in coeffs.coeffs().iter().enumerate().take(sp_0_deg + 1) {
unsafe {
*sp_coeffs_mat.get_unchecked_mut(j * width + i) = c_j;
}
}
}
let mut s_evals = dft.dft_batch(RowMajorMatrix::new(sp_coeffs_mat, width));
for (eq, row) in zip(eq_evals, s_evals.values.chunks_mut(width)) {
for x in row {
*x *= eq;
}
}
dft.idft_batch(s_evals)
};
let skip_domain_size = SC::F::from_usize(1 << l_skip);
let (numerator_term_per_air, denominator_term_per_air): (Vec<_>, Vec<_>) = (0..num_traces)
.map(|trace_idx| {
let [sum_claim_p, sum_claim_q] = [0, 1].map(|is_denom| {
(0..=s_0_deg)
.step_by(1 << l_skip)
.map(|j| unsafe {
*s_0_logup_polys
.values
.get_unchecked(j * 2 * num_traces + 2 * trace_idx + is_denom)
})
.sum::<SC::EF>()
* skip_domain_size
});
transcript.observe_ext(sum_claim_p);
transcript.observe_ext(sum_claim_q);
(sum_claim_p, sum_claim_q)
})
.unzip();
let mu = transcript.sample_ext();
debug!(%mu);
let mu_pows = mu.powers().take(3 * num_traces).collect_vec();
let s_0_zc_poly = {
let eq_uni = eq_uni_poly::<SC::F, _>(l_skip, prover.xi[0]);
let mut eq_coeffs = eq_uni.into_coeffs();
eq_coeffs.resize(large_uni_domain, SC::EF::ZERO);
let eq_evals = dft.dft(eq_coeffs);
let mut sp_coeffs = SC::EF::zero_vec(large_uni_domain);
let mus = &mu_pows[2 * num_traces..];
let polys = &sp_0_polys[2 * num_traces..];
for (j, batch_coeff) in sp_coeffs.iter_mut().enumerate().take(sp_0_deg + 1) {
for (&mu, poly) in zip(mus, polys) {
*batch_coeff += mu * *poly.coeffs().get(j).unwrap_or(&SC::EF::ZERO);
}
}
let mut s_evals = dft.dft(sp_coeffs);
for (eq, x) in zip(eq_evals, &mut s_evals) {
*x *= eq;
}
dft.idft(s_evals)
};
let s_0_poly = UnivariatePoly::new(
zip(
s_0_logup_polys.values.chunks_exact(2 * num_traces),
s_0_zc_poly,
)
.take(s_0_deg + 1)
.map(|(logup_row, batched_zc)| {
let coeff = batched_zc
+ zip(&mu_pows, logup_row)
.map(|(&mu_j, &x)| mu_j * x)
.sum::<SC::EF>();
transcript.observe_ext(coeff);
coeff
})
.collect(),
);
let r_0 = transcript.sample_ext();
r.push(r_0);
debug!(round = 0, r_round = %r_0);
prover.prev_s_eval = s_0_poly.eval_at_point(r_0);
debug!("s_0(r_0) = {}", prover.prev_s_eval);
prover.fold_ple_evals(ctx, r_0);
let _mle_rounds_span =
info_span!("prover.batch_constraints.mle_rounds", phase = "prover").entered();
debug!(%s_deg);
for round in 1..=n_max {
let sp_round_evals = prover.sumcheck_polys_eval(round, r[round - 1])?;
let tail_start = prover
.n_per_trace
.iter()
.find_position(|&&n| round as isize > n)
.map(|(i, _)| i)
.unwrap_or(num_traces);
let mut sp_head_zc = vec![SC::EF::ZERO; constraint_degree];
let mut sp_head_logup = vec![SC::EF::ZERO; constraint_degree];
let mut sp_tail = SC::EF::ZERO;
for trace_idx in 0..num_traces {
let zc_idx = 2 * num_traces + trace_idx;
let numer_idx = 2 * trace_idx;
let denom_idx = numer_idx + 1;
if trace_idx < tail_start {
for i in 0..constraint_degree {
sp_head_zc[i] += mu_pows[zc_idx] * sp_round_evals[zc_idx][i];
sp_head_logup[i] += mu_pows[numer_idx] * sp_round_evals[numer_idx][i]
+ mu_pows[denom_idx] * sp_round_evals[denom_idx][i];
}
} else {
sp_tail += mu_pows[zc_idx] * sp_round_evals[zc_idx][0]
+ mu_pows[numer_idx] * sp_round_evals[numer_idx][0]
+ mu_pows[denom_idx] * sp_round_evals[denom_idx][0];
}
}
let mut sp_head_evals = vec![SC::EF::ZERO; s_deg];
for i in 0..constraint_degree {
sp_head_evals[i + 1] = prover.eq_ns[round - 1] * sp_head_zc[i]
+ prover.eq_sharp_ns[round - 1] * sp_head_logup[i];
}
let xi_cur = prover.xi[l_skip + round - 1];
{
let eq_xi_0 = SC::EF::ONE - xi_cur;
let eq_xi_1 = xi_cur;
sp_head_evals[0] =
(prover.prev_s_eval - eq_xi_1 * sp_head_evals[1] - sp_tail) * eq_xi_0.inverse();
}
let sp_head = UnivariatePoly::lagrange_interpolate(
&(0..s_deg).map(SC::F::from_usize).collect_vec(),
&sp_head_evals,
);
let batch_s = {
let mut coeffs = sp_head.into_coeffs();
coeffs.push(SC::EF::ZERO);
let b = SC::EF::ONE - xi_cur;
let a = xi_cur - b;
for i in (0..s_deg).rev() {
coeffs[i + 1] = a * coeffs[i] + b * coeffs[i + 1];
}
coeffs[0] *= b;
coeffs[1] += sp_tail;
UnivariatePoly::new(coeffs)
};
let batch_s_evals = (1..=s_deg)
.map(|i| batch_s.eval_at_point(SC::EF::from_usize(i)))
.collect_vec();
for &eval in &batch_s_evals {
transcript.observe_ext(eval);
}
sumcheck_round_polys.push(batch_s_evals);
let r_round = transcript.sample_ext();
debug!(%round, %r_round);
r.push(r_round);
prover.prev_s_eval = batch_s.eval_at_point(r_round);
prover.fold_mle_evals(round, r_round);
}
drop(_mle_rounds_span);
if r.len() != n_max + 1 {
return Err(LogupZerocheckError::RLenMismatch {
r_len: r.len(),
expected: n_max + 1,
});
}
let column_openings = prover.into_column_openings()?;
for (helper, openings) in prover.eval_helpers.iter().zip(column_openings.iter()) {
for (claim, claim_rot) in column_openings_by_rot(&openings[0], helper.needs_next) {
transcript.observe_ext(claim);
transcript.observe_ext(claim_rot);
}
}
for (helper, openings) in prover.eval_helpers.iter().zip(column_openings.iter()) {
for part in openings.iter().skip(1) {
for (claim, claim_rot) in column_openings_by_rot(part, helper.needs_next) {
transcript.observe_ext(claim);
transcript.observe_ext(claim_rot);
}
}
}
let batch_constraint_proof = BatchConstraintProof::<SC> {
numerator_term_per_air,
denominator_term_per_air,
univariate_round_coeffs: s_0_poly.into_coeffs(),
sumcheck_round_polys,
column_openings,
};
let gkr_proof = GkrProof::<SC> {
logup_pow_witness,
q0_claim: frac_sum_proof.fractional_sum.1,
claims_per_layer: frac_sum_proof.claims_per_layer,
sumcheck_polys: frac_sum_proof.sumcheck_polys,
};
Ok((gkr_proof, batch_constraint_proof, r))
}