Skip to main content

openvm_stark_backend/prover/logup_zerocheck/
cpu.rs

1use std::{
2    cmp::max,
3    iter::{self, zip},
4    mem::take,
5};
6
7use itertools::Itertools;
8use p3_field::{ExtensionField, Field, PrimeCharacteristicRing, TwoAdicField};
9use p3_maybe_rayon::prelude::*;
10use p3_util::log2_strict_usize;
11
12use crate::{
13    air_builders::symbolic::{
14        symbolic_variable::Entry, SymbolicConstraints, SymbolicExpressionNode,
15    },
16    parizip,
17    poly_common::{eval_eq_mle, eval_eq_sharp_uni, eval_eq_uni, UnivariatePoly},
18    prover::{
19        error::LogupZerocheckError,
20        logup_zerocheck::EvalHelper,
21        poly::evals_eq_hypercubes,
22        stacked_pcs::StackedLayout,
23        sumcheck::{
24            batch_fold_mle_evals, batch_fold_ple_evals, fold_ple_evals, sumcheck_round0_deg,
25            sumcheck_round_poly_evals, sumcheck_uni_round0_poly,
26        },
27        ColMajorMatrix, CpuColMajorBackend, DeviceMultiStarkProvingKey, MatrixDimensions,
28        ProverBackend, ProvingContext, StridedColMajorMatrixView,
29    },
30    StarkProtocolConfig,
31};
32
33pub struct LogupZerocheckCpu<'a, SC: StarkProtocolConfig> {
34    pub alpha_logup: SC::EF,
35    pub beta_pows: Vec<SC::EF>,
36
37    pub l_skip: usize,
38    pub n_logup: usize,
39    pub n_max: usize,
40
41    pub omega_skip: SC::F,
42    pub omega_skip_pows: Vec<SC::F>,
43
44    pub interactions_layout: StackedLayout,
45    pub(crate) eval_helpers: Vec<EvalHelper<'a, SC::F>>,
46    /// Max constraint degree across constraints and interactions
47    pub constraint_degree: usize,
48    pub n_per_trace: Vec<isize>,
49    max_num_constraints: usize,
50
51    // Available after GKR:
52    pub xi: Vec<SC::EF>,
53    lambda_pows: Vec<SC::EF>,
54    // T -> segment tree of eq(xi[j..1+n_T]) for j=1..=n_T in _reverse_ layout
55    eq_xi_per_trace: Vec<Vec<SC::EF>>,
56    eq_3b_per_trace: Vec<Vec<SC::EF>>,
57    sels_per_trace_base: Vec<ColMajorMatrix<SC::F>>,
58    // After univariate round 0:
59    pub mat_evals_per_trace: Vec<Vec<ColMajorMatrix<SC::EF>>>,
60    pub sels_per_trace: Vec<ColMajorMatrix<SC::EF>>,
61    // Stores \hat{f}(\vec r_n) * r_{n+1} .. r_{round-1} for polys f that are "done" in the batch
62    // sumcheck
63    pub(crate) zerocheck_tilde_evals: Vec<SC::EF>,
64    pub(crate) logup_tilde_evals: Vec<[SC::EF; 2]>,
65
66    // In round `j`, contains `s_{j-1}(r_{j-1})`
67    pub(crate) prev_s_eval: SC::EF,
68    pub(crate) eq_ns: Vec<SC::EF>,
69    pub(crate) eq_sharp_ns: Vec<SC::EF>,
70}
71
72impl<'a, SC: StarkProtocolConfig> LogupZerocheckCpu<'a, SC>
73where
74    SC::F: TwoAdicField,
75    SC::EF: TwoAdicField + ExtensionField<SC::F>,
76    CpuColMajorBackend<SC>: ProverBackend<Val = SC::F, Matrix = ColMajorMatrix<SC::F>>,
77{
78    pub fn new(
79        pk: &'a DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
80        ctx: &ProvingContext<CpuColMajorBackend<SC>>,
81        n_logup: usize,
82        interactions_layout: StackedLayout,
83        alpha_logup: SC::EF,
84        beta_logup: SC::EF,
85    ) -> Result<Self, LogupZerocheckError> {
86        let l_skip = pk.params.l_skip;
87        let omega_skip = SC::F::two_adic_generator(l_skip);
88        let omega_skip_pows = omega_skip.powers().take(1 << l_skip).collect_vec();
89        let num_airs_present = ctx.per_trace.len();
90
91        let constraint_degree = pk.max_constraint_degree;
92        let max_interaction_length = ctx
93            .per_trace
94            .iter()
95            .flat_map(|(air_idx, _)| {
96                pk.per_air[*air_idx]
97                    .vk
98                    .symbolic_constraints
99                    .interactions
100                    .iter()
101                    .map(|i| i.message.len())
102            })
103            .max()
104            .unwrap_or(0);
105        let beta_pows = beta_logup
106            .powers()
107            .take(max_interaction_length + 1)
108            .collect_vec();
109
110        let n_per_trace: Vec<isize> = ctx
111            .common_main_traces()
112            .map(|(_, t)| log2_strict_usize(t.height()) as isize - l_skip as isize)
113            .collect_vec();
114        let n_max: usize = n_per_trace[0].max(0) as usize;
115
116        let eval_helpers: Vec<EvalHelper<SC::F>> = ctx
117            .per_trace
118            .iter()
119            .map(|(air_idx, trace_ctx)| {
120                let pk = &pk.per_air[*air_idx];
121                let constraints = &pk.vk.symbolic_constraints.constraints;
122                let public_values = trace_ctx.public_values.clone();
123                let preprocessed_trace: Option<StridedColMajorMatrixView<'_, SC::F>> = pk
124                    .preprocessed_data
125                    .as_ref()
126                    .map(|cd| cd.trace.as_view().into());
127                let partitioned_main_trace: Vec<StridedColMajorMatrixView<'_, SC::F>> = trace_ctx
128                    .cached_mains
129                    .iter()
130                    .map(|cd| cd.trace.as_view().into())
131                    .chain(iter::once(trace_ctx.common_main.as_view().into()))
132                    .collect_vec();
133                let constraint_degree = pk.vk.max_constraint_degree;
134                // Scan constraints to see if we need `next` row and also check index bounds
135                // so we don't need to check them per row.
136                let mut rotation = 0;
137                for node in &constraints.nodes {
138                    if let SymbolicExpressionNode::Variable(var) = node {
139                        match var.entry {
140                            Entry::Preprocessed { offset } => {
141                                rotation = max(rotation, offset);
142                                let width = preprocessed_trace.as_ref().unwrap().width();
143                                if var.index >= width {
144                                    return Err(
145                                        LogupZerocheckError::PreprocessedIndexOutOfBounds {
146                                            index: var.index,
147                                            width,
148                                        },
149                                    );
150                                }
151                            }
152                            Entry::Main { part_index, offset } => {
153                                rotation = max(rotation, offset);
154                                let width = partitioned_main_trace[part_index].width();
155                                if var.index >= width {
156                                    return Err(
157                                        LogupZerocheckError::MainPartitionIndexOutOfBounds {
158                                            part_index,
159                                            col_index: var.index,
160                                            width,
161                                        },
162                                    );
163                                }
164                            }
165                            Entry::Public => {
166                                if var.index >= public_values.len() {
167                                    return Err(LogupZerocheckError::PublicValueIndexOutOfBounds {
168                                        index: var.index,
169                                        len: public_values.len(),
170                                    });
171                                }
172                            }
173                            Entry::Challenge => {
174                                return Err(LogupZerocheckError::ChallengeNotSupported)
175                            }
176                        }
177                    }
178                }
179                let needs_next = pk.vk.params.need_rot;
180                debug_assert_eq!(needs_next, rotation > 0);
181                let symbolic_constraints = SymbolicConstraints::from(&pk.vk.symbolic_constraints);
182                Ok(EvalHelper {
183                    constraints_dag: &pk.vk.symbolic_constraints.constraints,
184                    interactions: symbolic_constraints.interactions,
185                    public_values,
186                    preprocessed_trace,
187                    needs_next,
188                    constraint_degree,
189                })
190            })
191            .collect::<Result<_, _>>()?; // end of preparation / loading of constraints
192        let max_num_constraints = pk
193            .per_air
194            .iter()
195            .map(|pk| pk.vk.symbolic_constraints.constraints.constraint_idx.len())
196            .max()
197            .unwrap_or(0);
198
199        let zerocheck_tilde_evals = vec![SC::EF::ZERO; num_airs_present];
200        let logup_tilde_evals = vec![[SC::EF::ZERO; 2]; num_airs_present];
201        Ok(Self {
202            alpha_logup,
203            beta_pows,
204            l_skip,
205            n_logup,
206            n_max,
207            omega_skip,
208            omega_skip_pows,
209            interactions_layout,
210            constraint_degree,
211            max_num_constraints,
212            n_per_trace,
213            eval_helpers,
214            xi: vec![],
215            lambda_pows: vec![],
216            sels_per_trace_base: vec![],
217            eq_xi_per_trace: vec![],
218            eq_3b_per_trace: vec![],
219            mat_evals_per_trace: vec![],
220            sels_per_trace: vec![],
221            zerocheck_tilde_evals,
222            logup_tilde_evals,
223            prev_s_eval: SC::EF::ZERO,
224            eq_ns: Vec::with_capacity(n_max + 1),
225            eq_sharp_ns: Vec::with_capacity(n_max + 1),
226        })
227    }
228
229    /// Returns the `s_0` polynomials in coefficient form. There should be exactly `num_airs_present
230    /// \* 3` polynomials, in the order `(s_0)_{p,T}, (s_0)_{q,T}, (s_0)_{zerocheck,T}` per trace
231    /// `T`. This is computed _before_ sampling batching randomness `mu` because the result is
232    /// used to observe the sum claims `sum_{p,T}, sum_{q,T}`. The `s_0` polynomials could be
233    /// returned in either coefficient or evaluation form, but we return them all in coefficient
234    /// form for uniformity and debugging since this interpolation is inexpensive.
235    pub fn sumcheck_uni_round0_polys(
236        &mut self,
237        ctx: &ProvingContext<CpuColMajorBackend<SC>>,
238        lambda: SC::EF,
239    ) -> Result<Vec<UnivariatePoly<SC::EF>>, LogupZerocheckError> {
240        let n_logup = self.n_logup;
241        let l_skip = self.l_skip;
242        let xi = &self.xi;
243        self.lambda_pows = lambda.powers().take(self.max_num_constraints).collect_vec();
244
245        // For each trace, for each interaction \hat\sigma, the eq(ΞΎ_3,b_{T,\hat\sigma}) term.
246        // This is some weight per interaction that does not depend on the row.
247        self.eq_3b_per_trace = self
248            .eval_helpers
249            .par_iter()
250            .zip(&self.n_per_trace)
251            .enumerate()
252            .map(
253                |(trace_idx, (helper, &n))| -> Result<Vec<SC::EF>, LogupZerocheckError> {
254                    // Everything for logup is done with respect to lifted traces
255                    // Note: `n_lift = \tilde{n}` from the paper
256                    let n_lift = n.max(0) as usize;
257                    if helper.interactions.is_empty() {
258                        return Ok(vec![]);
259                    }
260                    let mut b_vec = vec![SC::F::ZERO; n_logup - n_lift];
261                    (0..helper.interactions.len())
262                        .map(|i| {
263                            // PERF[jpw]: interactions_layout.get is linear
264                            let stacked_idx = self
265                                .interactions_layout
266                                .get(trace_idx, i)
267                                .ok_or(LogupZerocheckError::InteractionsLayoutMissing {
268                                    trace_idx,
269                                    interaction_idx: i,
270                                })?
271                                .row_idx;
272                            debug_assert!(stacked_idx.trailing_zeros() as usize >= n_lift + l_skip);
273                            let mut b_int = stacked_idx >> (l_skip + n_lift);
274                            for b in &mut b_vec {
275                                *b = SC::F::from_bool(b_int & 1 == 1);
276                                b_int >>= 1;
277                            }
278                            Ok(eval_eq_mle(&xi[l_skip + n_lift..l_skip + n_logup], &b_vec))
279                        })
280                        .collect()
281                },
282            )
283            .collect::<Result<Vec<_>, _>>()?;
284
285        // PERF[jpw]: make Hashmap from unique n -> eq_n(xi, -)
286        // NOTE: this is evaluations of `x -> eq_{H_{\tilde n}}(x, \xi[l_skip..l_skip + \tilde n])`
287        // on hypercube `H_{\tilde n}`. We store the univariate component eq_D separately as
288        // an optimization.
289        self.eq_xi_per_trace = self
290            .n_per_trace
291            .par_iter()
292            .map(|&n| {
293                let n_lift = n.max(0) as usize;
294                // PERF[jpw]: might be able to share computations between eq_xi, eq_sharp
295                // computations the eq(xi, -) evaluations on hyperprism for
296                // zerocheck
297                evals_eq_hypercubes(n_lift, xi[l_skip..l_skip + n_lift].iter().rev())
298            })
299            .collect();
300
301        // For each trace, create selectors as a 3-column matrix of _the lifts of_ [is_first,
302        // is_transition, is_last]
303        //
304        // PERF[jpw]: I think it's better to not save these and just
305        // interpolate directly using the formulas for selectors
306        self.sels_per_trace_base = self
307            .n_per_trace
308            .iter()
309            .map(|&n| {
310                let log_height = l_skip.checked_add_signed(n).unwrap();
311                let height = 1 << log_height;
312                let lifted_height = height.max(1 << l_skip);
313                let mut mat = SC::F::zero_vec(3 * lifted_height);
314                mat[lifted_height..2 * lifted_height].fill(SC::F::ONE);
315                for i in (0..lifted_height).step_by(height) {
316                    mat[i] = SC::F::ONE; // is_first
317                    mat[lifted_height + i + height - 1] = SC::F::ZERO; // is_transition
318                    mat[2 * lifted_height + i + height - 1] = SC::F::ONE; // is_last
319                }
320                ColMajorMatrix::new(mat, 3)
321            })
322            .collect_vec();
323
324        let sp_0_zerochecks = self
325            .eval_helpers
326            .par_iter()
327            .enumerate()
328            .map(|(trace_idx, helper)| {
329                let trace_ctx = &ctx.per_trace[trace_idx].1;
330                let n_lift = log2_strict_usize(trace_ctx.height()).saturating_sub(l_skip);
331                let mats = &helper.view_mats(trace_ctx);
332                let eq_xi = &self.eq_xi_per_trace[trace_idx][(1 << n_lift) - 1..(2 << n_lift) - 1];
333                let sels = self.sels_per_trace_base[trace_idx].as_view();
334                let mut parts = vec![(sels.into(), false)];
335                parts.extend_from_slice(mats);
336                // s'_0 has degree dependent on this AIR's constraint degree
337                // s'_0(Z) is a univariate polynomial which vanishes on D (zerocheck). Hence q(Z) =
338                // s'_0(Z) / Z_D(Z) = s'_0(Z) / (Z^{2^l_skip} - 1) is a polynomial of degree d *
339                // (2^l_skip - 1) - 2^l_skip = (d - 1) * 2^l_skip - d We can obtain
340                // q(Z) by interpolating evaluations on (d - 1) * 2^l_skip points. For computation
341                // efficiency, we choose these to be (d - 1) cosets of D. To avoid divide by zero,
342                // we avoid the coset equal to the subgroup D itself.
343                let constraint_deg = helper.constraint_degree as usize;
344                if constraint_deg == 0 {
345                    return UnivariatePoly(vec![]);
346                }
347                let num_cosets = constraint_deg - 1;
348                let [q] = sumcheck_uni_round0_poly(
349                    l_skip,
350                    n_lift,
351                    num_cosets,
352                    &parts,
353                    |z, x, row_parts| {
354                        let eq = eq_xi[x];
355                        let constraint_eval = helper.acc_constraints(row_parts, &self.lambda_pows);
356                        let zerofier = z.exp_power_of_2(l_skip) - SC::F::ONE;
357                        [eq * constraint_eval * zerofier.inverse()]
358                    },
359                );
360                // sp_0 = (Z^{2^l_skip} - 1) * q
361                let sp_0_deg = sumcheck_round0_deg(l_skip, constraint_deg);
362                let coeffs = (0..=sp_0_deg)
363                    .map(|i| {
364                        let mut c = -*q.coeffs().get(i).unwrap_or(&SC::EF::ZERO);
365                        if i >= 1 << l_skip {
366                            c += q.coeffs()[i - (1 << l_skip)];
367                        }
368                        c
369                    })
370                    .collect_vec();
371                debug_assert_eq!(
372                    coeffs.iter().step_by(1 << l_skip).copied().sum::<SC::EF>(),
373                    SC::EF::ZERO,
374                    "Zerocheck sum is not zero for air_id: {}",
375                    ctx.per_trace[trace_idx].0
376                );
377                UnivariatePoly(coeffs)
378            })
379            .collect::<Vec<_>>();
380        // Reminder: sum claims for zerocheck are zero, per AIR
381
382        // We interpolate each logup round 0 sumcheck poly because we need to use it to compute
383        // sum_{\hat{p}, T, I}, sum_{\hat{q}, T, I} per trace.
384        let sp_0_logups = self
385            .eval_helpers
386            .par_iter()
387            .enumerate()
388            .flat_map(|(trace_idx, helper)| {
389                if helper.interactions.is_empty() {
390                    return [(); 2].map(|_| UnivariatePoly::new(vec![]));
391                }
392                let trace_ctx = &ctx.per_trace[trace_idx].1;
393                let log_height = log2_strict_usize(trace_ctx.height());
394                let n_lift = log_height.saturating_sub(l_skip);
395                let mats = &helper.view_mats(trace_ctx);
396                let eq_xi = &self.eq_xi_per_trace[trace_idx][(1 << n_lift) - 1..(2 << n_lift) - 1];
397                let eq_3bs = &self.eq_3b_per_trace[trace_idx];
398                let sels = self.sels_per_trace_base[trace_idx].as_view();
399                let mut parts = vec![(sels.into(), false)];
400                parts.extend_from_slice(mats);
401                let norm_factor_denom = 1 << l_skip.saturating_sub(log_height);
402                let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
403
404                // degree is constraint_degree + 1 due to eq term
405                let [mut numer, denom] = sumcheck_uni_round0_poly(
406                    l_skip,
407                    n_lift,
408                    helper.constraint_degree as usize,
409                    &parts,
410                    |_z, x, row_parts| {
411                        let eq = eq_xi[x];
412                        let [numer, denom] =
413                            helper.acc_interactions(row_parts, &self.beta_pows, eq_3bs);
414                        [eq * numer, eq * denom]
415                    },
416                );
417                for p in numer.coeffs_mut() {
418                    *p *= norm_factor;
419                }
420                [numer, denom]
421            })
422            .collect::<Vec<_>>();
423
424        Ok(sp_0_logups.into_iter().chain(sp_0_zerochecks).collect())
425    }
426
427    /// After univariate sumcheck round 0, fold prismalinear evaluations using randomness `r_0`.
428    /// Folding _could_ directly mutate inplace the trace matrices in `ctx` as they will not be
429    /// needed after this.
430    pub fn fold_ple_evals(&mut self, ctx: &ProvingContext<CpuColMajorBackend<SC>>, r_0: SC::EF) {
431        let l_skip = self.l_skip;
432        // "Fold" all PLE evaluations by interpolating and evaluating at `r_0`.
433        // NOTE: after this folding, \hat{T} and \hat{T_{rot}} will be treated as completely
434        // distinct matrices.
435        self.mat_evals_per_trace = self
436            .eval_helpers
437            .par_iter()
438            .zip(ctx.per_trace.par_iter())
439            .map(|(helper, (_, trace_ctx))| {
440                let mats = helper.view_mats(trace_ctx);
441                mats.into_par_iter()
442                    .map(|(mat, is_rot)| fold_ple_evals(l_skip, mat, is_rot, r_0))
443                    .collect::<Vec<_>>()
444            })
445            .collect::<Vec<_>>();
446        self.sels_per_trace =
447            batch_fold_ple_evals(l_skip, take(&mut self.sels_per_trace_base), false, r_0);
448        let eq_r0 = eval_eq_uni(l_skip, self.xi[0], r_0);
449        let eq_sharp_r0 = eval_eq_sharp_uni(&self.omega_skip_pows, &self.xi[..l_skip], r_0);
450        self.eq_ns.push(eq_r0);
451        self.eq_sharp_ns.push(eq_sharp_r0);
452        self.eq_xi_per_trace.iter_mut().for_each(|eq| {
453            // trim the back (which corresponds to r_{j-1}) because we don't need it anymore
454            if eq.len() > 1 {
455                eq.truncate(eq.len() / 2);
456            }
457        });
458    }
459
460    /// Returns length `3 * num_airs_present` polynomials, each polynomial either evaluated at
461    /// `1,...,deg(s')` or at `1` if a linear term (terms in front-loaded sumcheck that have reached
462    /// exhaustion)
463    pub fn sumcheck_polys_eval(
464        &mut self,
465        round: usize,
466        r_prev: SC::EF,
467    ) -> Result<Vec<Vec<SC::EF>>, LogupZerocheckError> {
468        // sp = s'
469        let sp_deg = self.constraint_degree;
470        let eq_r_acc = *self
471            .eq_ns
472            .last()
473            .ok_or(LogupZerocheckError::EqNsEmpty { round })?;
474        let eq_sharp_r_acc = *self
475            .eq_sharp_ns
476            .last()
477            .ok_or(LogupZerocheckError::EqSharpNsEmpty { round })?;
478        let sp_zerocheck_evals: Vec<Vec<SC::EF>> = parizip!(
479            &self.eval_helpers,
480            &mut self.zerocheck_tilde_evals,
481            &self.n_per_trace,
482            &self.mat_evals_per_trace,
483            &self.sels_per_trace,
484            &self.eq_xi_per_trace
485        )
486        .map(|(helper, tilde_eval, &n, mats, sels, eq_xi_tree)| {
487            let n_lift = n.max(0) as usize;
488            if round > n_lift {
489                if round == n_lift + 1 {
490                    // Evaluate \hat{f}(\vec r_n)
491                    let parts = iter::once(sels)
492                        .chain(mats)
493                        .map(|mat| mat.columns().map(|c| c[0]).collect_vec())
494                        .collect_vec();
495                    // eq(xi, \vect r_{round-1})
496                    *tilde_eval = eq_r_acc * helper.acc_constraints(&parts, &self.lambda_pows);
497                } else {
498                    *tilde_eval *= r_prev;
499                };
500                vec![*tilde_eval]
501            } else {
502                let log_num_y = n_lift - round;
503                let num_y = 1 << log_num_y;
504                let eq_xi = &eq_xi_tree[num_y - 1..];
505                let parts = iter::once(sels)
506                    .chain(mats)
507                    .map(|m| m.as_view())
508                    .collect_vec();
509                let [s] =
510                    sumcheck_round_poly_evals(log_num_y + 1, sp_deg, &parts, |_x, y, row_parts| {
511                        let eq = eq_xi[y];
512                        let constraint_eval = helper.acc_constraints(row_parts, &self.lambda_pows);
513                        [eq * constraint_eval]
514                    });
515                s
516            }
517        })
518        .collect();
519
520        let sp_logup_evals: Vec<Vec<SC::EF>> = parizip!(
521            &self.eval_helpers,
522            &mut self.logup_tilde_evals,
523            &self.n_per_trace,
524            &self.mat_evals_per_trace,
525            &self.sels_per_trace,
526            &self.eq_xi_per_trace,
527            &self.eq_3b_per_trace
528        )
529        .flat_map(|(helper, tilde_eval, &n, mats, sels, eq_xi_tree, eq_3bs)| {
530            if helper.interactions.is_empty() {
531                return [vec![SC::EF::ZERO; sp_deg], vec![SC::EF::ZERO; sp_deg]];
532            }
533            let n_lift = n.max(0) as usize;
534            let norm_factor_denom = 1 << (-n).max(0);
535            let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
536            if round > n_lift {
537                if round == n_lift + 1 {
538                    // Evaluate \hat{f}(\vec r_n)
539                    let parts = iter::once(sels)
540                        .chain(mats)
541                        .map(|mat| mat.columns().map(|c| c[0]).collect_vec())
542                        .collect_vec();
543                    *tilde_eval = helper
544                        .acc_interactions(&parts, &self.beta_pows, eq_3bs)
545                        .map(|x| eq_sharp_r_acc * x);
546                    tilde_eval[0] *= norm_factor;
547                } else {
548                    for x in tilde_eval.iter_mut() {
549                        *x *= r_prev;
550                    }
551                };
552                tilde_eval.map(|tilde_eval| vec![tilde_eval])
553            } else {
554                let parts = iter::once(sels)
555                    .chain(mats)
556                    .map(|m| m.as_view())
557                    .collect_vec();
558                let log_num_y = n_lift - round;
559                let num_y = 1 << log_num_y;
560                let eq_xi = &eq_xi_tree[num_y - 1..];
561                let [mut numer, denom] =
562                    sumcheck_round_poly_evals(log_num_y + 1, sp_deg, &parts, |_x, y, row_parts| {
563                        let eq = eq_xi[y];
564                        helper
565                            .acc_interactions(row_parts, &self.beta_pows, eq_3bs)
566                            .map(|eval| eq * eval)
567                    });
568                for p in &mut numer {
569                    *p *= norm_factor;
570                }
571                [numer, denom]
572            }
573        })
574        .collect();
575
576        Ok(sp_logup_evals
577            .into_iter()
578            .chain(sp_zerocheck_evals)
579            .collect())
580    }
581
582    pub fn fold_mle_evals(&mut self, round: usize, r_round: SC::EF) {
583        self.mat_evals_per_trace = take(&mut self.mat_evals_per_trace)
584            .into_iter()
585            .map(|mats| batch_fold_mle_evals(mats, r_round))
586            .collect_vec();
587        self.sels_per_trace = batch_fold_mle_evals(take(&mut self.sels_per_trace), r_round);
588        self.eq_xi_per_trace.par_iter_mut().for_each(|eq| {
589            // trim the back (which corresponds to r_{j-1}) because we don't need it anymore
590            if eq.len() > 1 {
591                eq.truncate(eq.len() / 2);
592            }
593        });
594        let xi = self.xi[self.l_skip + round - 1];
595        let eq_r = eval_eq_mle(&[xi], &[r_round]);
596        self.eq_ns.push(self.eq_ns[round - 1] * eq_r);
597        self.eq_sharp_ns.push(self.eq_sharp_ns[round - 1] * eq_r);
598
599        #[allow(unused_variables)]
600        #[cfg(debug_assertions)]
601        if tracing::enabled!(tracing::Level::DEBUG) && round == self.n_max {
602            use itertools::izip;
603
604            for (trace_idx, (helper, &n, mats, sels, eq_xi)) in izip!(
605                &self.eval_helpers,
606                &self.n_per_trace,
607                &self.mat_evals_per_trace,
608                &self.sels_per_trace,
609                &self.eq_xi_per_trace
610            )
611            .enumerate()
612            {
613                let parts = iter::once(sels)
614                    .chain(mats)
615                    .map(|mat| mat.columns().map(|c| c[0]).collect_vec())
616                    .collect_vec();
617                let expr = helper.acc_constraints(&parts, &self.lambda_pows);
618                tracing::debug!(%trace_idx, %expr, "constraints_eval");
619            }
620
621            for (trace_idx, (helper, &n, mats, sels, eq_3bs)) in izip!(
622                &self.eval_helpers,
623                &self.n_per_trace,
624                &self.mat_evals_per_trace,
625                &self.sels_per_trace,
626                &self.eq_3b_per_trace
627            )
628            .enumerate()
629            {
630                if helper.interactions.is_empty() {
631                    continue;
632                }
633                let parts = iter::once(sels)
634                    .chain(mats)
635                    .map(|mat| mat.columns().map(|c| c[0]).collect_vec())
636                    .collect_vec();
637                let [num, denom] = helper.acc_interactions(&parts, &self.beta_pows, eq_3bs);
638
639                tracing::debug!(%trace_idx, %num, %denom, "interactions_eval");
640            }
641        }
642    }
643
644    pub fn into_column_openings(&mut self) -> Result<Vec<Vec<Vec<SC::EF>>>, LogupZerocheckError> {
645        let num_airs_present = self.mat_evals_per_trace.len();
646        let mut column_openings = Vec::with_capacity(num_airs_present);
647        // At the end, we've folded all MLEs so they only have one row equal to evaluation at `\vec
648        // r`.
649        for (helper, mut mat_evals) in self
650            .eval_helpers
651            .iter()
652            .zip(take(&mut self.mat_evals_per_trace))
653        {
654            // For column openings, we pop common_main (and common_main_rot when present) and put it
655            // at the front.
656            let openings_of_air: Vec<Vec<SC::EF>> = if helper.needs_next {
657                let common_main_rot = mat_evals
658                    .pop()
659                    .ok_or(LogupZerocheckError::MatEvalsPopNone)?;
660                let common_main = mat_evals
661                    .pop()
662                    .ok_or(LogupZerocheckError::MatEvalsPopNone)?;
663                iter::once(&[common_main, common_main_rot] as &[_])
664                    .chain(mat_evals.chunks_exact(2))
665                    .map(|pair| {
666                        zip(pair[0].columns(), pair[1].columns())
667                            .flat_map(|(claim, claim_rot)| {
668                                debug_assert_eq!(claim.len(), 1);
669                                debug_assert_eq!(claim_rot.len(), 1);
670                                [claim[0], claim_rot[0]]
671                            })
672                            .collect_vec()
673                    })
674                    .collect_vec()
675            } else {
676                let common_main = mat_evals
677                    .pop()
678                    .ok_or(LogupZerocheckError::MatEvalsPopNone)?;
679                iter::once(common_main)
680                    .chain(mat_evals.into_iter())
681                    .map(|mat| {
682                        mat.columns()
683                            .map(|claim| {
684                                debug_assert_eq!(claim.len(), 1);
685                                claim[0]
686                            })
687                            .collect_vec()
688                    })
689                    .collect_vec()
690            };
691            column_openings.push(openings_of_air);
692        }
693        Ok(column_openings)
694    }
695}