use std::{array::from_fn, collections::HashMap, iter::zip, mem::take};
use itertools::Itertools;
use p3_field::{ExtensionField, PrimeCharacteristicRing, TwoAdicField};
use p3_maybe_rayon::prelude::*;
use tracing::{debug, instrument};
use crate::{
poly_common::{eval_eq_mle, eval_eq_uni, eval_eq_uni_at_one, eval_in_uni, UnivariatePoly},
proof::StackingProof,
prover::{
poly::evals_eq_hypercube,
stacked_pcs::{StackedPcsData, StackedSlice},
sumcheck::{
batch_fold_mle_evals, fold_mle_evals, fold_ple_evals, sumcheck_round0_deg,
sumcheck_round_poly_evals, sumcheck_uni_round0_poly,
},
ColMajorMatrix, ColMajorMatrixView, CpuColMajorBackend, MatrixDimensions, MatrixView,
ProverBackend, ReferenceDevice,
},
FiatShamirTranscript, StarkProtocolConfig,
};
pub trait StackedReductionProver<'a, PB: ProverBackend, PD> {
fn new(
device: &'a PD,
stacked_per_commit: Vec<&'a PB::PcsData>,
need_rot_per_commit: Vec<Vec<bool>>,
r: &[PB::Challenge],
lambda: PB::Challenge,
) -> Self;
fn batch_sumcheck_uni_round0_poly(&mut self) -> UnivariatePoly<PB::Challenge>;
fn fold_ple_evals(&mut self, u_0: PB::Challenge);
fn batch_sumcheck_poly_eval(
&mut self,
round: usize,
u_prev: PB::Challenge,
) -> [PB::Challenge; 2];
fn fold_mle_evals(&mut self, round: usize, u_round: PB::Challenge);
fn into_stacked_openings(self) -> Vec<Vec<PB::Challenge>>;
}
#[instrument(level = "info", skip_all)]
pub fn prove_stacked_opening_reduction<'a, SC, PB, PD, TS, SRP>(
device: &'a PD,
transcript: &mut TS,
n_stack: usize,
stacked_per_commit: Vec<&'a PB::PcsData>,
need_rot_per_commit: Vec<Vec<bool>>,
r: &[PB::Challenge],
) -> (StackingProof<SC>, Vec<PB::Challenge>)
where
SC: StarkProtocolConfig,
PB: ProverBackend<Val = SC::F, Challenge = SC::EF>,
TS: FiatShamirTranscript<SC>,
SRP: StackedReductionProver<'a, PB, PD>,
{
let lambda = transcript.sample_ext();
let mut prover = SRP::new(device, stacked_per_commit, need_rot_per_commit, r, lambda);
let s_0 = prover.batch_sumcheck_uni_round0_poly();
for &coeff in s_0.coeffs() {
transcript.observe_ext(coeff);
}
let mut u_vec = Vec::with_capacity(n_stack + 1);
let u_0 = transcript.sample_ext();
u_vec.push(u_0);
debug!(round = 0, u_round = %u_0);
prover.fold_ple_evals(u_0);
let mut sumcheck_round_polys = Vec::with_capacity(n_stack);
#[allow(clippy::needless_range_loop)]
for round in 1..=n_stack {
let batch_s_evals = prover.batch_sumcheck_poly_eval(round, u_vec[round - 1]);
for &eval in &batch_s_evals {
transcript.observe_ext(eval);
}
sumcheck_round_polys.push(batch_s_evals);
let u_round = transcript.sample_ext();
u_vec.push(u_round);
debug!(%round, %u_round);
prover.fold_mle_evals(round, u_round);
}
let stacking_openings = prover.into_stacked_openings();
for claims_for_com in &stacking_openings {
for &claim in claims_for_com {
transcript.observe_ext(claim);
}
}
let proof = StackingProof::<SC> {
univariate_round_coeffs: s_0.0,
sumcheck_round_polys,
stacking_openings,
};
(proof, u_vec)
}
pub struct StackedReductionCpu<'a, SC: StarkProtocolConfig> {
l_skip: usize,
omega_skip: SC::F,
r_0: SC::EF,
lambda_pows: Vec<SC::EF>,
eq_const: SC::EF,
stacked_per_commit: Vec<&'a StackedPcsData<SC::F, SC::Digest>>,
trace_views: Vec<TraceViewMeta>,
ht_diff_idxs: Vec<usize>,
eq_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>>,
k_rot_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>>,
q_evals: Vec<ColMajorMatrix<SC::EF>>,
eq_ub_per_trace: Vec<SC::EF>,
}
struct TraceViewMeta {
com_idx: usize,
slice: StackedSlice,
lambda_eq_idx: usize,
lambda_rot_idx: Option<usize>,
}
impl<'a, SC: StarkProtocolConfig>
StackedReductionProver<'a, CpuColMajorBackend<SC>, ReferenceDevice<SC>>
for StackedReductionCpu<'a, SC>
where
SC::F: TwoAdicField,
SC::EF: TwoAdicField + ExtensionField<SC::F>,
CpuColMajorBackend<SC>: ProverBackend<
Val = SC::F,
Challenge = SC::EF,
PcsData = StackedPcsData<SC::F, SC::Digest>,
Matrix = ColMajorMatrix<SC::F>,
>,
{
fn new(
device: &ReferenceDevice<SC>,
stacked_per_commit: Vec<&'a StackedPcsData<SC::F, SC::Digest>>,
need_rot_per_commit: Vec<Vec<bool>>,
r: &[SC::EF],
lambda: SC::EF,
) -> Self {
let l_skip = device.params().l_skip;
let omega_skip = SC::F::two_adic_generator(l_skip);
let mut trace_views = Vec::new();
let mut lambda_idx = 0usize;
for (com_idx, d) in stacked_per_commit.iter().enumerate() {
let need_rot_for_commit = &need_rot_per_commit[com_idx];
debug_assert_eq!(need_rot_for_commit.len(), d.layout.mat_starts.len());
for &(mat_idx, _col_idx, slice) in &d.layout.sorted_cols {
let lambda_eq_idx = lambda_idx;
lambda_idx += 1;
let lambda_rot_idx = if need_rot_for_commit[mat_idx] {
Some(lambda_idx)
} else {
None
};
lambda_idx += 1;
trace_views.push(TraceViewMeta {
com_idx,
slice,
lambda_eq_idx,
lambda_rot_idx,
});
}
}
let lambda_pows = lambda.powers().take(lambda_idx).collect_vec();
let mut ht_diff_idxs = Vec::new();
let mut eq_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>> = HashMap::new();
let mut last_height = 0;
for (i, tv) in trace_views.iter().enumerate() {
let n_lift = tv.slice.log_height().saturating_sub(l_skip);
if i == 0 || tv.slice.log_height() != last_height {
ht_diff_idxs.push(i);
last_height = tv.slice.log_height();
}
eq_r_per_lht
.entry(tv.slice.log_height())
.or_insert_with(|| ColMajorMatrix::new(evals_eq_hypercube(&r[1..1 + n_lift]), 1));
}
ht_diff_idxs.push(trace_views.len());
let eq_const = eval_eq_uni_at_one(l_skip, r[0] * omega_skip);
let eq_ub_per_trace = vec![SC::EF::ONE; trace_views.len()];
Self {
l_skip,
omega_skip,
r_0: r[0],
lambda_pows,
eq_const,
stacked_per_commit,
trace_views,
ht_diff_idxs,
eq_r_per_lht,
q_evals: vec![],
k_rot_r_per_lht: HashMap::new(),
eq_ub_per_trace,
}
}
fn batch_sumcheck_uni_round0_poly(&mut self) -> UnivariatePoly<SC::EF> {
let l_skip = self.l_skip;
let omega_skip = self.omega_skip;
let s_0_deg = sumcheck_round0_deg(l_skip, 2);
let s_0_polys: Vec<_> = self
.ht_diff_idxs
.par_windows(2)
.flat_map(|window| {
let t_window = &self.trace_views[window[0]..window[1]];
let log_height = t_window[0].slice.log_height();
let n = log_height as isize - l_skip as isize;
let n_lift = n.max(0) as usize;
let eq_rs = self.eq_r_per_lht.get(&log_height).unwrap().column(0);
debug_assert_eq!(eq_rs.len(), 1 << n_lift);
let q_t_cols = t_window
.iter()
.map(|tv| {
debug_assert_eq!(tv.slice.log_height(), log_height);
let q = &self.stacked_per_commit[tv.com_idx].matrix;
let s = tv.slice;
let q_t_col = &q.column(s.col_idx)[s.row_idx..s.row_idx + s.len(l_skip)];
(ColMajorMatrixView::new(q_t_col, 1).into(), false)
})
.collect_vec();
sumcheck_uni_round0_poly(l_skip, n_lift, 2, &q_t_cols, |z, x, evals| {
let eq_cube = eq_rs[x];
let (l, omega, r_uni) = if n.is_negative() {
(
l_skip.wrapping_add_signed(n),
omega_skip.exp_power_of_2(-n as usize),
self.r_0.exp_power_of_2(-n as usize),
)
} else {
(l_skip, omega_skip, self.r_0)
};
let ind = eval_in_uni(l_skip, n, z);
let eq_uni_r0 = eval_eq_uni(l, z.into(), r_uni);
let eq_uni_r0_rot = eval_eq_uni(l, z.into(), r_uni * omega);
let eq_uni_1 = eval_eq_uni_at_one(l_skip, z);
let k_rot_cube = eq_rs[rot_prev(x, n_lift)];
let eq = eq_uni_r0 * eq_cube;
let k_rot =
eq_uni_r0_rot * eq_cube + self.eq_const * eq_uni_1 * (k_rot_cube - eq_cube);
zip(t_window, evals).fold([SC::EF::ZERO; 2], |mut acc, (tv, eval)| {
let q = eval[0];
acc[0] += self.lambda_pows[tv.lambda_eq_idx] * eq * q * ind;
if let Some(rot_idx) = tv.lambda_rot_idx {
acc[1] += self.lambda_pows[rot_idx] * k_rot * q * ind;
}
acc
})
})
})
.collect();
let s_0_coeffs = (0..=s_0_deg)
.map(|i| {
s_0_polys
.iter()
.map(|evals| evals.coeffs()[i])
.sum::<SC::EF>()
})
.collect_vec();
UnivariatePoly::new(s_0_coeffs)
}
fn fold_ple_evals(&mut self, u_0: SC::EF) {
let l_skip = self.l_skip;
let r_0 = self.r_0;
let omega_skip = self.omega_skip;
self.q_evals = self
.stacked_per_commit
.iter()
.map(|d| fold_ple_evals(l_skip, d.matrix.as_view().into(), false, u_0))
.collect_vec();
let eq_uni_u0r0 = eval_eq_uni(l_skip, u_0, r_0);
let eq_uni_u0r0_rot = eval_eq_uni(l_skip, u_0, r_0 * omega_skip);
let eq_uni_u01 = eval_eq_uni_at_one(l_skip, u_0);
self.k_rot_r_per_lht = self
.eq_r_per_lht
.par_iter_mut()
.map(|(&log_height, mat)| {
let n = log_height as isize - l_skip as isize;
let n_lift = n.max(0) as usize;
debug_assert_eq!(mat.values.len(), 1 << n_lift);
let ind = eval_in_uni(l_skip, n, u_0);
let (eq_uni, eq_uni_rot) = if n.is_negative() {
let omega = omega_skip.exp_power_of_2(-n as usize);
let r = r_0.exp_power_of_2(-n as usize);
let l = l_skip.wrapping_add_signed(n);
(eval_eq_uni(l, u_0, r), eval_eq_uni(l, u_0, r * omega))
} else {
(eq_uni_u0r0, eq_uni_u0r0_rot)
};
let evals: Vec<_> = (0..1 << n_lift)
.into_par_iter()
.map(|x| {
let eq_cube = unsafe { *mat.get_unchecked(x, 0) };
let k_rot_cube = unsafe { *mat.get_unchecked(rot_prev(x, n_lift), 0) };
ind * (eq_uni_rot * eq_cube
+ self.eq_const * eq_uni_u01 * (k_rot_cube - eq_cube))
})
.collect();
mat.values.par_iter_mut().for_each(|v| {
*v *= ind * eq_uni;
});
(log_height, ColMajorMatrix::new(evals, 1))
})
.collect();
}
fn batch_sumcheck_poly_eval(&mut self, round: usize, _u_prev: SC::EF) -> [SC::EF; 2] {
let l_skip = self.l_skip;
let s_deg = 2;
let s_evals: Vec<_> = self
.ht_diff_idxs
.par_windows(2)
.flat_map(|window| {
let t_views = &self.trace_views[window[0]..window[1]];
let log_height = t_views[0].slice.log_height();
let n_lift = log_height.saturating_sub(l_skip); let hypercube_dim = n_lift.saturating_sub(round);
let eq_rs = self.eq_r_per_lht.get(&log_height).unwrap().column(0);
let k_rot_rs = self.k_rot_r_per_lht.get(&log_height).unwrap().column(0);
debug_assert_eq!(eq_rs.len(), 1 << n_lift.saturating_sub(round - 1));
debug_assert_eq!(k_rot_rs.len(), 1 << n_lift.saturating_sub(round - 1));
let t_cols = t_views
.iter()
.map(|tv| {
debug_assert_eq!(tv.slice.log_height(), log_height);
let q = &self.q_evals[tv.com_idx];
let s = tv.slice;
let row_start = if round <= n_lift {
(s.row_idx >> log_height) << (hypercube_dim + 1)
} else {
(s.row_idx >> (l_skip + round)) << 1
};
let t_col =
&q.column(s.col_idx)[row_start..row_start + (2 << hypercube_dim)];
ColMajorMatrixView::new(t_col, 1)
})
.collect_vec();
sumcheck_round_poly_evals(hypercube_dim + 1, s_deg, &t_cols, |x, y, evals| {
evals
.iter()
.enumerate()
.fold([SC::EF::ZERO; 2], |mut acc, (i, eval)| {
let t_idx = window[0] + i;
let tv = &self.trace_views[t_idx];
let q = eval[0];
let mut eq_ub = self.eq_ub_per_trace[t_idx];
let (eq, k_rot) = if round > n_lift {
let b = (tv.slice.row_idx >> (l_skip + round - 1)) & 1;
eq_ub *= eval_eq_mle(&[x], &[SC::F::from_bool(b == 1)]);
debug_assert_eq!(y, 0);
(eq_rs[0] * eq_ub, k_rot_rs[0] * eq_ub)
} else {
let eq_r =
eq_rs[y << 1] * (SC::EF::ONE - x) + eq_rs[(y << 1) + 1] * x;
let k_rot_r = k_rot_rs[y << 1] * (SC::EF::ONE - x)
+ k_rot_rs[(y << 1) + 1] * x;
(eq_r * eq_ub, k_rot_r * eq_ub)
};
acc[0] += self.lambda_pows[tv.lambda_eq_idx] * q * eq;
if let Some(rot_idx) = tv.lambda_rot_idx {
acc[1] += self.lambda_pows[rot_idx] * q * k_rot;
}
acc
})
})
})
.collect();
from_fn(|i| s_evals.iter().map(|evals| evals[i]).sum::<SC::EF>())
}
fn fold_mle_evals(&mut self, round: usize, u_round: SC::EF) {
let l_skip = self.l_skip;
self.q_evals = batch_fold_mle_evals(take(&mut self.q_evals), u_round);
self.eq_r_per_lht = take(&mut self.eq_r_per_lht)
.into_par_iter()
.map(|(lht, mat)| (lht, fold_mle_evals(mat, u_round)))
.collect();
self.k_rot_r_per_lht = take(&mut self.k_rot_r_per_lht)
.into_par_iter()
.map(|(lht, mat)| (lht, fold_mle_evals(mat, u_round)))
.collect();
for (tv, eq_ub) in zip(&self.trace_views, &mut self.eq_ub_per_trace) {
let s = tv.slice;
let n_lift = s.log_height().saturating_sub(l_skip);
if round > n_lift {
let b = (s.row_idx >> (l_skip + round - 1)) & 1;
*eq_ub *= eval_eq_mle(&[u_round], &[SC::F::from_bool(b == 1)]);
}
}
}
fn into_stacked_openings(self) -> Vec<Vec<SC::EF>> {
self.q_evals
.into_iter()
.map(|q| {
debug_assert_eq!(q.height(), 1);
q.values
})
.collect()
}
}
fn rot_prev(x_int: usize, n: usize) -> usize {
debug_assert!(x_int < (1 << n));
if x_int == 0 {
(1 << n) - 1
} else {
x_int - 1
}
}