Skip to main content

openvm_stark_backend/prover/
stacked_reduction.rs

1//! Stacked opening reduction
2
3use 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
26/// Helper trait for proving the reduction of column opening claims and column rotation opening
27/// claims to opening claims of column polynomials of the stacked matrix.
28///
29/// Returns the reduction proof and the random vector `u` of length `1 + n_stack`.
30pub trait StackedReductionProver<'a, PB: ProverBackend, PD> {
31    /// We only provide a view to the stacked `PcsData` per commitment because the WHIR prover will
32    /// still use the PLE evaluations of the stacked matrices later. The order of
33    /// `stacked_per_commit` is `common_main, preprocessed for trace_idx=0 (if any), cached_0 for
34    /// trace_idx=0, ..., preprocessed for trace_idx=1 (if any), ...`.
35    ///
36    /// The `lambda` is the batching randomness for the batch sumcheck.
37    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    /// Return the `s_0` batched polynomial from univariate round 0 of sumcheck.
46    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/// Batch sumcheck to reduce trace openings, including rotations, to stacked matrix opening.
62///
63/// The `stacked_matrix, stacked_layout` should be the result of stacking the `traces` with
64/// parameters `l_skip` and `n_stack`.
65#[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    // Batching randomness
81    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    // end round 0
96
97    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    // After round 0:
143    k_rot_r_per_lht: HashMap<usize, ColMajorMatrix<SC::EF>>,
144    q_evals: Vec<ColMajorMatrix<SC::EF>>,
145    /// Stores eq(u[1+n_T..round-1], b_{T,j}[..round-n_T-1])
146    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        // +1 from eq term
241        let s_0_deg = sumcheck_round0_deg(l_skip, 2);
242        // We want to compute algebraic batching, via \lambda,
243        // for each (T, j) pair of (trace, column) of the univariate polynomials
244        // ```text
245        // Z -> sum_{x in H_{n_stack}} q(Z,x) in_{D,n_T}(Z) eq_{D_{n_T}}((Z,x[..\tilde n_T]), r[..1+\tilde n_T]) eq(x[\tilde n_T..], b_{T,j})
246        // Z -> sum_{x in H_{n_stack}} q(Z,x) in_{D,n_T}(Z) \kappa_{\rot, D_{n_T}}((Z,x[..\tilde n_T]), r[..1+\tilde n_T]) eq(x[\tilde n_T..], b_{T,j})
247        // ```
248        // where `b_{T,j}` is length `n_stack - n_T` binary encoding of `StackedSlice.row_idx >>
249        // (l_skip
250        // + n_T)`. Note that since x is in the hypercube, by definition of `eq` the above
251        // simplifies to
252        // ```text
253        // Z -> sum_{x in H_{n_T}} q(Z,x,b_{T,j}) (in_{D,n_T}(Z) eq(Z,r_0)) eq(x[..\tilde n_T], r[1..1+\tilde n_T])
254        // Z -> sum_{x in H_{n_T}} q(Z,x,b_{T,j}) in_{D,n_T}(Z) \kappa_rot((Z, x[..n_T]), (r_0, r[1..1+n_T]))
255        // ```
256        // where we also simplified the other `eq` term.
257        // We further simplify the second using equation
258        // ```text
259        // \kappa_rot((Z, x[..n_T]), (r_0, r[1..1+n_T])) =
260        // eq_D(Z,omega_D r_0) eq(x[..n_T], r[1..1+n_T]) + eq_D(Z,1)eq_D(omega_D r_0,1) ( kappa_rot(x[..n_T], r[1..1+n_T]) - eq(x[..n_T], r[1..1+n_T]) )
261        // ```
262        // We compute the last polynomial in our usual way, by considering `q(\vec Z, b_{T,j})` as a
263        // prismalinear polynomial and using its evaluations on `D_{n_T}`.
264        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                // Prepare the q subslice eval views
275                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                        // NOTE: even if s.stride(l_skip) != 1, we use the full non-strided column
283                        // subslice. The sumcheck will not depend on the values outside of the
284                        // stride because of the `in_{D, n_T}` indicator below.
285                        (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                    // eq_uni_1, k_rot_cube are only used when n > 0
303                    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        // fold PLEs into MLEs for \eq and \kappa_\rot, using u_0
341        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        // \kappa_\rot(x, r) = eq(rot^{-1}(x), r)
345        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                // folded \kappa_\rot evals
362                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                // update \eq with the univariate factor:
372                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        // We want to compute algebraic batching, via \lambda,
384        // for each (T, j) pair of (trace, column) of the univariate polynomials
385        // ```
386        // X -> sum_{y in H_{n_stack-round}} q(u[..round],X,y) eq((u[..round],X,y[..n_T-round]), r[..1+n_T]) eq(y[n_T-round..], b_{T,j})
387        //      = sum_{y in H_{n_T-round}} q(u[..round],X,y,b_{T,j}) eq((u[..round],X,y), r[..1+n_T])
388        // X -> sum_{y in H_{n_stack-round}} q(u[..round],X,y) \kappa_\rot((u[..round],X,y[..n_T]), r[..1+n_T]) eq(y[n_T..], b_{T,j})
389        // ```
390        // if `round <= n_T`. Otherwise we compute
391        // ```
392        // X -> sum_{y in H_{n_stack-round}} q(u[..round],X,y) eq((u[..1+n_T], r[..1+n_T]) eq((u[1+n_T..round],X,y[round..]), b_{T,j})
393        //      = q(u[..round], X, b_{T,j}[round-n_T..]) eq((u[..1+n_T], r[..1+n_T]) eq((u[1+n_T..round],X), b_{T,j}[..round-n_T])
394        // X -> sum_{y in H_{n_stack-round}} q(u[..round],X,y) \kappa_\rot(u[..1+n_T], r[..1+n_T]) eq((u[1+n_T..round],X,y[round..]), b_{T,j})
395        // ```
396        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); // \tilde{n}_T
403                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                // Prepare the q subslice eval views
409                let t_cols = t_views
410                    .iter()
411                    .map(|tv| {
412                        debug_assert_eq!(tv.slice.log_height(), log_height);
413                        // q(u[..round], X, b_{T,j}[round-\tilde n_T..])
414                        // q_evals has been folded already
415                        let q = &self.q_evals[tv.com_idx];
416                        let s = tv.slice;
417                        let row_start = if round <= n_lift {
418                            // round >= 1 so n_lift >= 1
419                            (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                                // Extra contribution of eq(X, b_{T,j}[round-n_T-1])
439                                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                                // linearly interpolate eq(-, r[..1+n_T]), \kappa_\rot(-,
445                                // r[..1+n_T])
446                                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                // Folding above did nothing, and we update the eq(u[1+n_T..=round],
480                // b_{T,j}[..=round-n_T-1]) value
481                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
498/// `x_int` is the integer representation of point on H_n.
499fn 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}