1use std::{array::from_fn, collections::HashMap, iter::zip, mem::take};
4
5use itertools::Itertools;
6use p3_field::{ExtensionField, PrimeCharacteristicRing, TwoAdicField};
7use p3_maybe_rayon::prelude::*;
8use tracing::{debug, instrument};
9
10use crate::{
11 poly_common::{eval_eq_mle, eval_eq_uni, eval_eq_uni_at_one, eval_in_uni, UnivariatePoly},
12 proof::StackingProof,
13 prover::{
14 poly::evals_eq_hypercube,
15 stacked_pcs::{StackedPcsData, StackedSlice},
16 sumcheck::{
17 batch_fold_mle_evals, fold_mle_evals, fold_ple_evals, sumcheck_round0_deg,
18 sumcheck_round_poly_evals, sumcheck_uni_round0_poly,
19 },
20 ColMajorMatrix, ColMajorMatrixView, CpuColMajorBackend, MatrixDimensions, MatrixView,
21 ProverBackend, ReferenceDevice,
22 },
23 FiatShamirTranscript, StarkProtocolConfig,
24};
25
26pub trait StackedReductionProver<'a, PB: ProverBackend, PD> {
31 fn new(
38 device: &'a PD,
39 stacked_per_commit: Vec<&'a PB::PcsData>,
40 need_rot_per_commit: Vec<Vec<bool>>,
41 r: &[PB::Challenge],
42 lambda: PB::Challenge,
43 ) -> Self;
44
45 fn batch_sumcheck_uni_round0_poly(&mut self) -> UnivariatePoly<PB::Challenge>;
47
48 fn fold_ple_evals(&mut self, u_0: PB::Challenge);
49
50 fn batch_sumcheck_poly_eval(
51 &mut self,
52 round: usize,
53 u_prev: PB::Challenge,
54 ) -> [PB::Challenge; 2];
55
56 fn fold_mle_evals(&mut self, round: usize, u_round: PB::Challenge);
57
58 fn into_stacked_openings(self) -> Vec<Vec<PB::Challenge>>;
59}
60
61#[instrument(level = "info", skip_all)]
66pub fn prove_stacked_opening_reduction<'a, SC, PB, PD, TS, SRP>(
67 device: &'a PD,
68 transcript: &mut TS,
69 n_stack: usize,
70 stacked_per_commit: Vec<&'a PB::PcsData>,
71 need_rot_per_commit: Vec<Vec<bool>>,
72 r: &[PB::Challenge],
73) -> (StackingProof<SC>, Vec<PB::Challenge>)
74where
75 SC: StarkProtocolConfig,
76 PB: ProverBackend<Val = SC::F, Challenge = SC::EF>,
77 TS: FiatShamirTranscript<SC>,
78 SRP: StackedReductionProver<'a, PB, PD>,
79{
80 let lambda = transcript.sample_ext();
82
83 let mut prover = SRP::new(device, stacked_per_commit, need_rot_per_commit, r, lambda);
84 let s_0 = prover.batch_sumcheck_uni_round0_poly();
85 for &coeff in s_0.coeffs() {
86 transcript.observe_ext(coeff);
87 }
88
89 let mut u_vec = Vec::with_capacity(n_stack + 1);
90 let u_0 = transcript.sample_ext();
91 u_vec.push(u_0);
92 debug!(round = 0, u_round = %u_0);
93
94 prover.fold_ple_evals(u_0);
95 let mut sumcheck_round_polys = Vec::with_capacity(n_stack);
98
99 #[allow(clippy::needless_range_loop)]
100 for round in 1..=n_stack {
101 let batch_s_evals = prover.batch_sumcheck_poly_eval(round, u_vec[round - 1]);
102
103 for &eval in &batch_s_evals {
104 transcript.observe_ext(eval);
105 }
106 sumcheck_round_polys.push(batch_s_evals);
107
108 let u_round = transcript.sample_ext();
109 u_vec.push(u_round);
110 debug!(%round, %u_round);
111
112 prover.fold_mle_evals(round, u_round);
113 }
114 let stacking_openings = prover.into_stacked_openings();
115 for claims_for_com in &stacking_openings {
116 for &claim in claims_for_com {
117 transcript.observe_ext(claim);
118 }
119 }
120 let proof = StackingProof::<SC> {
121 univariate_round_coeffs: s_0.0,
122 sumcheck_round_polys,
123 stacking_openings,
124 };
125 (proof, u_vec)
126}
127
128pub struct StackedReductionCpu<'a, SC: StarkProtocolConfig> {
129 l_skip: usize,
130 omega_skip: SC::F,
131
132 r_0: SC::EF,
133 lambda_pows: Vec<SC::EF>,
134 eq_const: SC::EF,
135
136 stacked_per_commit: Vec<&'a StackedPcsData<SC::F, SC::Digest>>,
137 trace_views: Vec<TraceViewMeta>,
138 ht_diff_idxs: Vec<usize>,
139
140 eq_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>>,
141
142 k_rot_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>>,
144 q_evals: Vec<ColMajorMatrix<SC::EF>>,
145 eq_ub_per_trace: Vec<SC::EF>,
147}
148
149struct TraceViewMeta {
150 com_idx: usize,
151 slice: StackedSlice,
152 lambda_eq_idx: usize,
153 lambda_rot_idx: Option<usize>,
154}
155
156impl<'a, SC: StarkProtocolConfig>
157 StackedReductionProver<'a, CpuColMajorBackend<SC>, ReferenceDevice<SC>>
158 for StackedReductionCpu<'a, SC>
159where
160 SC::F: TwoAdicField,
161 SC::EF: TwoAdicField + ExtensionField<SC::F>,
162 CpuColMajorBackend<SC>: ProverBackend<
163 Val = SC::F,
164 Challenge = SC::EF,
165 PcsData = StackedPcsData<SC::F, SC::Digest>,
166 Matrix = ColMajorMatrix<SC::F>,
167 >,
168{
169 fn new(
170 device: &ReferenceDevice<SC>,
171 stacked_per_commit: Vec<&'a StackedPcsData<SC::F, SC::Digest>>,
172 need_rot_per_commit: Vec<Vec<bool>>,
173 r: &[SC::EF],
174 lambda: SC::EF,
175 ) -> Self {
176 let l_skip = device.params().l_skip;
177 let omega_skip = SC::F::two_adic_generator(l_skip);
178
179 let mut trace_views = Vec::new();
180 let mut lambda_idx = 0usize;
181 for (com_idx, d) in stacked_per_commit.iter().enumerate() {
182 let need_rot_for_commit = &need_rot_per_commit[com_idx];
183 debug_assert_eq!(need_rot_for_commit.len(), d.layout.mat_starts.len());
184 for &(mat_idx, _col_idx, slice) in &d.layout.sorted_cols {
185 let lambda_eq_idx = lambda_idx;
186 lambda_idx += 1;
187 let lambda_rot_idx = if need_rot_for_commit[mat_idx] {
188 Some(lambda_idx)
189 } else {
190 None
191 };
192 lambda_idx += 1;
193 trace_views.push(TraceViewMeta {
194 com_idx,
195 slice,
196 lambda_eq_idx,
197 lambda_rot_idx,
198 });
199 }
200 }
201 let lambda_pows = lambda.powers().take(lambda_idx).collect_vec();
202
203 let mut ht_diff_idxs = Vec::new();
204 let mut eq_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>> = HashMap::new();
205 let mut last_height = 0;
206 for (i, tv) in trace_views.iter().enumerate() {
207 let n_lift = tv.slice.log_height().saturating_sub(l_skip);
208 if i == 0 || tv.slice.log_height() != last_height {
209 ht_diff_idxs.push(i);
210 last_height = tv.slice.log_height();
211 }
212 eq_r_per_lht
213 .entry(tv.slice.log_height())
214 .or_insert_with(|| ColMajorMatrix::new(evals_eq_hypercube(&r[1..1 + n_lift]), 1));
215 }
216 ht_diff_idxs.push(trace_views.len());
217
218 let eq_const = eval_eq_uni_at_one(l_skip, r[0] * omega_skip);
219 let eq_ub_per_trace = vec![SC::EF::ONE; trace_views.len()];
220
221 Self {
222 l_skip,
223 omega_skip,
224 r_0: r[0],
225 lambda_pows,
226 eq_const,
227 stacked_per_commit,
228 trace_views,
229 ht_diff_idxs,
230 eq_r_per_lht,
231 q_evals: vec![],
232 k_rot_r_per_lht: HashMap::new(),
233 eq_ub_per_trace,
234 }
235 }
236
237 fn batch_sumcheck_uni_round0_poly(&mut self) -> UnivariatePoly<SC::EF> {
238 let l_skip = self.l_skip;
239 let omega_skip = self.omega_skip;
240 let s_0_deg = sumcheck_round0_deg(l_skip, 2);
242 let s_0_polys: Vec<_> = self
265 .ht_diff_idxs
266 .par_windows(2)
267 .flat_map(|window| {
268 let t_window = &self.trace_views[window[0]..window[1]];
269 let log_height = t_window[0].slice.log_height();
270 let n = log_height as isize - l_skip as isize;
271 let n_lift = n.max(0) as usize;
272 let eq_rs = self.eq_r_per_lht.get(&log_height).unwrap().column(0);
273 debug_assert_eq!(eq_rs.len(), 1 << n_lift);
274 let q_t_cols = t_window
276 .iter()
277 .map(|tv| {
278 debug_assert_eq!(tv.slice.log_height(), log_height);
279 let q = &self.stacked_per_commit[tv.com_idx].matrix;
280 let s = tv.slice;
281 let q_t_col = &q.column(s.col_idx)[s.row_idx..s.row_idx + s.len(l_skip)];
282 (ColMajorMatrixView::new(q_t_col, 1).into(), false)
286 })
287 .collect_vec();
288 sumcheck_uni_round0_poly(l_skip, n_lift, 2, &q_t_cols, |z, x, evals| {
289 let eq_cube = eq_rs[x];
290 let (l, omega, r_uni) = if n.is_negative() {
291 (
292 l_skip.wrapping_add_signed(n),
293 omega_skip.exp_power_of_2(-n as usize),
294 self.r_0.exp_power_of_2(-n as usize),
295 )
296 } else {
297 (l_skip, omega_skip, self.r_0)
298 };
299 let ind = eval_in_uni(l_skip, n, z);
300 let eq_uni_r0 = eval_eq_uni(l, z.into(), r_uni);
301 let eq_uni_r0_rot = eval_eq_uni(l, z.into(), r_uni * omega);
302 let eq_uni_1 = eval_eq_uni_at_one(l_skip, z);
304 let k_rot_cube = eq_rs[rot_prev(x, n_lift)];
305
306 let eq = eq_uni_r0 * eq_cube;
307 let k_rot =
308 eq_uni_r0_rot * eq_cube + self.eq_const * eq_uni_1 * (k_rot_cube - eq_cube);
309 zip(t_window, evals).fold([SC::EF::ZERO; 2], |mut acc, (tv, eval)| {
310 let q = eval[0];
311 acc[0] += self.lambda_pows[tv.lambda_eq_idx] * eq * q * ind;
312 if let Some(rot_idx) = tv.lambda_rot_idx {
313 acc[1] += self.lambda_pows[rot_idx] * k_rot * q * ind;
314 }
315 acc
316 })
317 })
318 })
319 .collect();
320 let s_0_coeffs = (0..=s_0_deg)
321 .map(|i| {
322 s_0_polys
323 .iter()
324 .map(|evals| evals.coeffs()[i])
325 .sum::<SC::EF>()
326 })
327 .collect_vec();
328 UnivariatePoly::new(s_0_coeffs)
329 }
330
331 fn fold_ple_evals(&mut self, u_0: SC::EF) {
332 let l_skip = self.l_skip;
333 let r_0 = self.r_0;
334 let omega_skip = self.omega_skip;
335 self.q_evals = self
336 .stacked_per_commit
337 .iter()
338 .map(|d| fold_ple_evals(l_skip, d.matrix.as_view().into(), false, u_0))
339 .collect_vec();
340 let eq_uni_u0r0 = eval_eq_uni(l_skip, u_0, r_0);
342 let eq_uni_u0r0_rot = eval_eq_uni(l_skip, u_0, r_0 * omega_skip);
343 let eq_uni_u01 = eval_eq_uni_at_one(l_skip, u_0);
344 self.k_rot_r_per_lht = self
346 .eq_r_per_lht
347 .par_iter_mut()
348 .map(|(&log_height, mat)| {
349 let n = log_height as isize - l_skip as isize;
350 let n_lift = n.max(0) as usize;
351 debug_assert_eq!(mat.values.len(), 1 << n_lift);
352 let ind = eval_in_uni(l_skip, n, u_0);
353 let (eq_uni, eq_uni_rot) = if n.is_negative() {
354 let omega = omega_skip.exp_power_of_2(-n as usize);
355 let r = r_0.exp_power_of_2(-n as usize);
356 let l = l_skip.wrapping_add_signed(n);
357 (eval_eq_uni(l, u_0, r), eval_eq_uni(l, u_0, r * omega))
358 } else {
359 (eq_uni_u0r0, eq_uni_u0r0_rot)
360 };
361 let evals: Vec<_> = (0..1 << n_lift)
363 .into_par_iter()
364 .map(|x| {
365 let eq_cube = unsafe { *mat.get_unchecked(x, 0) };
366 let k_rot_cube = unsafe { *mat.get_unchecked(rot_prev(x, n_lift), 0) };
367 ind * (eq_uni_rot * eq_cube
368 + self.eq_const * eq_uni_u01 * (k_rot_cube - eq_cube))
369 })
370 .collect();
371 mat.values.par_iter_mut().for_each(|v| {
373 *v *= ind * eq_uni;
374 });
375 (log_height, ColMajorMatrix::new(evals, 1))
376 })
377 .collect();
378 }
379
380 fn batch_sumcheck_poly_eval(&mut self, round: usize, _u_prev: SC::EF) -> [SC::EF; 2] {
381 let l_skip = self.l_skip;
382 let s_deg = 2;
383 let s_evals: Vec<_> = self
397 .ht_diff_idxs
398 .par_windows(2)
399 .flat_map(|window| {
400 let t_views = &self.trace_views[window[0]..window[1]];
401 let log_height = t_views[0].slice.log_height();
402 let n_lift = log_height.saturating_sub(l_skip); let hypercube_dim = n_lift.saturating_sub(round);
404 let eq_rs = self.eq_r_per_lht.get(&log_height).unwrap().column(0);
405 let k_rot_rs = self.k_rot_r_per_lht.get(&log_height).unwrap().column(0);
406 debug_assert_eq!(eq_rs.len(), 1 << n_lift.saturating_sub(round - 1));
407 debug_assert_eq!(k_rot_rs.len(), 1 << n_lift.saturating_sub(round - 1));
408 let t_cols = t_views
410 .iter()
411 .map(|tv| {
412 debug_assert_eq!(tv.slice.log_height(), log_height);
413 let q = &self.q_evals[tv.com_idx];
416 let s = tv.slice;
417 let row_start = if round <= n_lift {
418 (s.row_idx >> log_height) << (hypercube_dim + 1)
420 } else {
421 (s.row_idx >> (l_skip + round)) << 1
422 };
423 let t_col =
424 &q.column(s.col_idx)[row_start..row_start + (2 << hypercube_dim)];
425 ColMajorMatrixView::new(t_col, 1)
426 })
427 .collect_vec();
428 sumcheck_round_poly_evals(hypercube_dim + 1, s_deg, &t_cols, |x, y, evals| {
429 evals
430 .iter()
431 .enumerate()
432 .fold([SC::EF::ZERO; 2], |mut acc, (i, eval)| {
433 let t_idx = window[0] + i;
434 let tv = &self.trace_views[t_idx];
435 let q = eval[0];
436 let mut eq_ub = self.eq_ub_per_trace[t_idx];
437 let (eq, k_rot) = if round > n_lift {
438 let b = (tv.slice.row_idx >> (l_skip + round - 1)) & 1;
440 eq_ub *= eval_eq_mle(&[x], &[SC::F::from_bool(b == 1)]);
441 debug_assert_eq!(y, 0);
442 (eq_rs[0] * eq_ub, k_rot_rs[0] * eq_ub)
443 } else {
444 let eq_r =
447 eq_rs[y << 1] * (SC::EF::ONE - x) + eq_rs[(y << 1) + 1] * x;
448 let k_rot_r = k_rot_rs[y << 1] * (SC::EF::ONE - x)
449 + k_rot_rs[(y << 1) + 1] * x;
450 (eq_r * eq_ub, k_rot_r * eq_ub)
451 };
452 acc[0] += self.lambda_pows[tv.lambda_eq_idx] * q * eq;
453 if let Some(rot_idx) = tv.lambda_rot_idx {
454 acc[1] += self.lambda_pows[rot_idx] * q * k_rot;
455 }
456 acc
457 })
458 })
459 })
460 .collect();
461 from_fn(|i| s_evals.iter().map(|evals| evals[i]).sum::<SC::EF>())
462 }
463
464 fn fold_mle_evals(&mut self, round: usize, u_round: SC::EF) {
465 let l_skip = self.l_skip;
466 self.q_evals = batch_fold_mle_evals(take(&mut self.q_evals), u_round);
467 self.eq_r_per_lht = take(&mut self.eq_r_per_lht)
468 .into_par_iter()
469 .map(|(lht, mat)| (lht, fold_mle_evals(mat, u_round)))
470 .collect();
471 self.k_rot_r_per_lht = take(&mut self.k_rot_r_per_lht)
472 .into_par_iter()
473 .map(|(lht, mat)| (lht, fold_mle_evals(mat, u_round)))
474 .collect();
475 for (tv, eq_ub) in zip(&self.trace_views, &mut self.eq_ub_per_trace) {
476 let s = tv.slice;
477 let n_lift = s.log_height().saturating_sub(l_skip);
478 if round > n_lift {
479 let b = (s.row_idx >> (l_skip + round - 1)) & 1;
482 *eq_ub *= eval_eq_mle(&[u_round], &[SC::F::from_bool(b == 1)]);
483 }
484 }
485 }
486
487 fn into_stacked_openings(self) -> Vec<Vec<SC::EF>> {
488 self.q_evals
489 .into_iter()
490 .map(|q| {
491 debug_assert_eq!(q.height(), 1);
492 q.values
493 })
494 .collect()
495 }
496}
497
498fn rot_prev(x_int: usize, n: usize) -> usize {
500 debug_assert!(x_int < (1 << n));
501 if x_int == 0 {
502 (1 << n) - 1
503 } else {
504 x_int - 1
505 }
506}