Skip to main content

openvm_cpu_backend/logup_zerocheck/
mod.rs

1//! Row-major LogupZerocheck implementation.
2//!
3//! Optimizations over the reference backend:
4//! 1. Eliminates full row-major → col-major conversion for round 0 sumcheck
5//! 2. Uses batch DFT on extracted row-major blocks, leveraging SIMD through plonky3's butterfly
6//!    operations (PackedField in apply_to_rows)
7//! 3. Direct row-major access for constraint/interaction evaluation
8
9use std::{
10    cmp::max,
11    iter::{self, zip},
12    mem::take,
13    ops::{Add, Mul, Neg, Sub},
14};
15
16use itertools::{izip, Itertools};
17use openvm_stark_backend::{
18    air_builders::symbolic::{
19        symbolic_expression::SymbolicEvaluator,
20        symbolic_variable::{Entry, SymbolicVariable},
21        SymbolicConstraints, SymbolicExpressionDag, SymbolicExpressionNode,
22    },
23    calculate_n_logup,
24    dft::Radix2BowersSerial,
25    interaction::SymbolicInteraction,
26    poly_common::{eq_uni_poly, eval_eq_mle, eval_eq_sharp_uni, eval_eq_uni, UnivariatePoly},
27    proof::{column_openings_by_rot, BatchConstraintProof, GkrProof},
28    prover::{
29        error::LogupZerocheckError,
30        fractional_sumcheck_gkr::{fractional_sumcheck, Frac},
31        poly::{eq_sharp_uni_poly, evals_eq_hypercubes},
32        stacked_pcs::StackedLayout,
33        sumcheck::sumcheck_round0_deg,
34        AirProvingContext, DeviceMultiStarkProvingKey, MatrixDimensions, ProverBackend,
35        ProvingContext,
36    },
37    FiatShamirTranscript, StarkProtocolConfig,
38};
39use p3_dft::TwoAdicSubgroupDft;
40use p3_field::{
41    batch_multiplicative_inverse, ExtensionField, Field, PackedValue, PrimeCharacteristicRing,
42    TwoAdicField,
43};
44use p3_matrix::dense::RowMajorMatrix;
45use p3_maybe_rayon::prelude::*;
46use p3_util::log2_strict_usize;
47use tracing::{debug, info_span, instrument};
48
49use crate::backend::CpuBackend;
50
51// ============================================================================
52// Batch DFT helpers for row-major sumcheck
53// ============================================================================
54
55/// Extract `2^l_skip` rows from a row-major matrix for hypercube point `x`.
56/// Returns a RowMajorMatrix suitable for batch DFT.
57///
58/// For the hyperprism D_n evaluation, the rows for point `x` are:
59///   row[(x << l_skip) + z + offset] for z = 0..2^l_skip
60/// where offset=1 for rotation, 0 otherwise.
61#[inline]
62fn extract_rm_block<F: TwoAdicField>(
63    rm: &RowMajorMatrix<F>,
64    x: usize,
65    l_skip: usize,
66    offset: usize,
67) -> RowMajorMatrix<F> {
68    let height = rm.values.len() / rm.width;
69    let w = rm.width;
70    let sz = 1usize << l_skip;
71    let base = x << l_skip;
72    let mut vals = Vec::with_capacity(sz * w);
73    for z in 0..sz {
74        let r = (base + z + offset) % height;
75        let start = r * w;
76        vals.extend_from_slice(&rm.values[start..start + w]);
77    }
78    RowMajorMatrix::new(vals, w)
79}
80
81/// Batch coset DFT: given iDFT'd coefficients in a RowMajorMatrix,
82/// evaluate on the coset `shift * D` where D = <omega_skip>.
83///
84/// Steps: 1. Multiply row i by shift^i (twisting)
85///        2. Forward batch DFT
86#[inline]
87fn batch_coset_dft<F: TwoAdicField>(coeffs: &RowMajorMatrix<F>, shift: F) -> RowMajorMatrix<F> {
88    let w = coeffs.width;
89    let mut mat = coeffs.clone();
90    // Twist: multiply row i by shift^i
91    let mut s = F::ONE;
92    for chunk in mat.values.chunks_exact_mut(w) {
93        if s != F::ONE {
94            for v in chunk.iter_mut() {
95                *v *= s;
96            }
97        }
98        s *= shift;
99    }
100    Radix2BowersSerial.dft_batch(mat)
101}
102
103/// Extract blocks from selector/trace matrices, iDFT them, and compute coset evaluations.
104///
105/// Shared between `sumcheck_uni_round0_batch` and `sumcheck_uni_round0_zerocheck_packed`.
106fn extract_and_dft_blocks<F: TwoAdicField>(
107    x: usize,
108    l_skip: usize,
109    rm_mats: &[(&RowMajorMatrix<F>, bool)],
110    sels_rm: &RowMajorMatrix<F>,
111    coset_shifts: &[F],
112) -> (Vec<RowMajorMatrix<F>>, Vec<Vec<RowMajorMatrix<F>>>) {
113    let dft = Radix2BowersSerial;
114
115    let sels_block = extract_rm_block(sels_rm, x, l_skip, 0);
116    let sels_coeffs = dft.idft_batch(sels_block);
117    let sels_cosets = coset_shifts
118        .iter()
119        .map(|&shift| batch_coset_dft(&sels_coeffs, shift))
120        .collect();
121
122    let mat_coeffs: Vec<RowMajorMatrix<F>> = rm_mats
123        .iter()
124        .map(|(rm, is_rot)| {
125            let block = extract_rm_block(rm, x, l_skip, usize::from(*is_rot));
126            dft.idft_batch(block)
127        })
128        .collect();
129    let mat_cosets = mat_coeffs
130        .iter()
131        .map(|coeffs| {
132            coset_shifts
133                .iter()
134                .map(|&shift| batch_coset_dft(coeffs, shift))
135                .collect()
136        })
137        .collect();
138
139    (sels_cosets, mat_cosets)
140}
141
142/// Fold PLE evaluations directly from row-major data, avoiding the O(n*m) transpose.
143///
144/// Barycentric interpolation: for each column j,
145///   result[j] = scaling_factor * sum_i (col_scale[i] * mat[row_i, j])
146///
147/// Since rows are contiguous in row-major layout, this is a weighted sum of rows
148/// — extremely cache-friendly compared to column-major interpolation.
149fn fold_ple_evals_rowmajor<F, EF>(
150    l_skip: usize,
151    rm: &RowMajorMatrix<F>,
152    is_rot: bool,
153    r: EF,
154) -> RowMajorMatrix<EF>
155where
156    F: TwoAdicField,
157    EF: ExtensionField<F> + TwoAdicField,
158{
159    let height = rm.values.len() / rm.width;
160    let width = rm.width;
161    let lifted_height = height.max(1 << l_skip);
162    let skip_sz = 1usize << l_skip;
163    let new_height = lifted_height >> l_skip;
164    let offset = usize::from(is_rot);
165
166    // Precompute barycentric weights (same as in p3-interpolation)
167    let omega = F::two_adic_generator(l_skip);
168    let omega_pows: Vec<F> = omega.powers().take(skip_sz).collect_vec();
169    let denoms: Vec<EF> = omega_pows
170        .iter()
171        .map(|&x_i| r - EF::from(x_i))
172        .collect_vec();
173    let inv_denoms = batch_multiplicative_inverse(&denoms);
174
175    // col_scale[i] = omega^i / (r - omega^i)
176    let col_scale: Vec<EF> = omega_pows
177        .iter()
178        .zip(&inv_denoms)
179        .map(|(&sg, &diff_inv)| diff_inv * sg)
180        .collect_vec();
181
182    // scaling_factor = (r^N - 1) / N where N = 2^l_skip
183    let r_pow_n = r.exp_power_of_2(l_skip);
184    let scaling_factor = (r_pow_n - EF::ONE) * EF::from_usize(skip_sz).inverse();
185
186    // Output is row-major: values[x * width + col] for each (x, col)
187    let values: Vec<EF> = (0..new_height)
188        .into_par_iter()
189        .flat_map(|x| {
190            // For this x-point, compute weighted sum of 2^l_skip rows
191            let mut result = vec![EF::ZERO; width];
192            for z in 0..skip_sz {
193                let row_idx = ((x << l_skip) + z + offset) % height;
194                let row_start = row_idx * width;
195                let row = &rm.values[row_start..row_start + width];
196                let w = col_scale[z];
197                for (j, &val) in row.iter().enumerate() {
198                    result[j] += w * val;
199                }
200            }
201            // Apply scaling factor
202            for v in &mut result {
203                *v *= scaling_factor;
204            }
205            result
206        })
207        .collect();
208
209    RowMajorMatrix::new(values, width)
210}
211
212/// Run the round-0 sumcheck using batch DFT from row-major data.
213///
214/// This replaces the pattern of: to_col_major_owned → sumcheck_uni_round0_poly
215/// with direct row-major block extraction → batch iDFT → batch coset DFT → callback.
216///
217/// Key optimizations:
218/// 1. No full matrix transpose (saves ~100-200ms allocation + O(n) work)
219/// 2. Batch DFT processes all columns at once; butterfly `apply_to_rows` uses PackedField SIMD
220/// 3. Row extraction from row-major is contiguous (cache-friendly)
221fn sumcheck_uni_round0_batch<F, EF, FN, const WD: usize>(
222    l_skip: usize,
223    n: usize,
224    d: usize,
225    rm_mats: &[(&RowMajorMatrix<F>, bool)],
226    sels_rm: &RowMajorMatrix<F>,
227    w: FN,
228) -> [UnivariatePoly<EF>; WD]
229where
230    F: TwoAdicField,
231    EF: ExtensionField<F> + TwoAdicField,
232    FN: Fn(F, usize, &[Vec<F>]) -> [EF; WD] + Sync,
233{
234    if d == 0 {
235        return std::array::from_fn(|_| UnivariatePoly::new(vec![]));
236    }
237    let g = F::GENERATOR;
238    let omega_skip = F::two_adic_generator(l_skip);
239    let coset_shifts: Vec<F> = g.powers().skip(1).take(d).collect_vec();
240    let skip_sz = 1usize << l_skip;
241
242    // Map-Reduce over x ∈ H_n
243    let evals = (0..1usize << n).into_par_iter().map(|x| {
244        let (sels_cosets, mat_cosets) =
245            extract_and_dft_blocks(x, l_skip, rm_mats, sels_rm, &coset_shifts);
246
247        // Pre-allocate row_parts buffers: reused across all z-points to avoid
248        // per-z-point allocation (saves ~1.4GB of allocations for keccakf)
249        let sels_w = sels_cosets[0].width;
250        let mut row_parts: Vec<Vec<F>> = Vec::with_capacity(1 + mat_cosets.len());
251        row_parts.push(vec![F::ZERO; sels_w]);
252        for mc in &mat_cosets {
253            row_parts.push(vec![F::ZERO; mc[0].width]);
254        }
255
256        // Evaluate callback at each z-point across all cosets
257        let mut results = Vec::with_capacity(d * skip_sz);
258        for (z_idx, z) in omega_skip.powers().take(skip_sz).enumerate() {
259            for (ci, &shift) in coset_shifts.iter().enumerate() {
260                // Copy into pre-allocated buffers (no new allocation)
261                let ss = z_idx * sels_w;
262                row_parts[0].copy_from_slice(&sels_cosets[ci].values[ss..ss + sels_w]);
263
264                for (mat_idx, mc) in mat_cosets.iter().enumerate() {
265                    let mw = mc[ci].width;
266                    let ms = z_idx * mw;
267                    row_parts[1 + mat_idx].copy_from_slice(&mc[ci].values[ms..ms + mw]);
268                }
269
270                results.push(w(shift * z, x, &row_parts));
271            }
272        }
273        results
274    });
275
276    // Reduce: sum over H_n
277    let hypercube_sum = |mut acc: Vec<[EF; WD]>, x: Vec<[EF; WD]>| {
278        for (acc_i, x_i) in acc.iter_mut().zip(x) {
279            for (a, b) in acc_i.iter_mut().zip(x_i) {
280                *a += b;
281            }
282        }
283        acc
284    };
285    cfg_if::cfg_if! {
286        if #[cfg(feature = "parallel")] {
287            let evals = evals.reduce(
288                || vec![[EF::ZERO; WD]; d << l_skip],
289                hypercube_sum,
290            );
291        } else {
292            let evals: Vec<_> = evals.collect();
293            let evals = evals.into_iter().fold(
294                vec![[EF::ZERO; WD]; d << l_skip],
295                hypercube_sum,
296            );
297        }
298    }
299
300    // Assemble polynomial from coset evaluations
301    std::array::from_fn(|i| {
302        let values: Vec<EF> = evals.iter().map(|x| x[i]).collect_vec();
303        UnivariatePoly::from_geometric_cosets_evals_idft(RowMajorMatrix::new(values, d), g, g)
304    })
305}
306
307/// Packed SIMD zerocheck round 0: evaluates the constraint DAG for WIDTH
308/// z-points simultaneously using F::Packing, reducing DAG walks by WIDTH×.
309///
310/// On aarch64/NEON (WIDTH=4): reduces 32 DAG walks to 8 per x-point.
311/// On x86/AVX2 (WIDTH=8): reduces 32 DAG walks to 4 per x-point.
312///
313/// The zerofier `(shift*z)^(2^l_skip) - 1 = shift^(2^l_skip) - 1` is
314/// constant per coset (since z^(2^l_skip) = 1 for z in D_{l_skip}),
315/// so we precompute zerofier_inv per coset.
316fn sumcheck_uni_round0_zerocheck_packed<SC>(
317    l_skip: usize,
318    n: usize,
319    d: usize,
320    rm_mats: &[(&RowMajorMatrix<SC::F>, bool)],
321    sels_rm: &RowMajorMatrix<SC::F>,
322    helper: &RowMajorEvalHelper<'_, SC>,
323    eq_xi: &[SC::EF],
324    lambda_pows: &[SC::EF],
325) -> UnivariatePoly<SC::EF>
326where
327    SC: StarkProtocolConfig,
328    SC::F: TwoAdicField,
329    SC::EF: TwoAdicField + ExtensionField<SC::F>,
330{
331    if d == 0 {
332        return UnivariatePoly::new(vec![]);
333    }
334    let g = SC::F::GENERATOR;
335    let coset_shifts: Vec<SC::F> = g.powers().skip(1).take(d).collect_vec();
336    let skip_sz = 1usize << l_skip;
337    let width = <SC::F as Field>::Packing::WIDTH;
338
339    // Precompute zerofier_inv per coset: shift^(2^l_skip) - 1 is constant per coset
340    let zerofier_invs: Vec<SC::F> = coset_shifts
341        .iter()
342        .map(|&shift| (shift.exp_power_of_2(l_skip) - SC::F::ONE).inverse())
343        .collect();
344
345    // Map-Reduce over x ∈ H_n
346    let evals = (0..1usize << n).into_par_iter().map(|x| {
347        let eq = eq_xi[x];
348
349        let (sels_cosets, mat_cosets) =
350            extract_and_dft_blocks(x, l_skip, rm_mats, sels_rm, &coset_shifts);
351
352        // Pre-allocate packed row_parts buffers
353        let sels_w = sels_cosets[0].width;
354        let mut packed_row_parts: Vec<Vec<<SC::F as Field>::Packing>> =
355            Vec::with_capacity(1 + mat_cosets.len());
356        packed_row_parts.push(vec![<SC::F as Field>::Packing::default(); sels_w]);
357        for mc in &mat_cosets {
358            packed_row_parts.push(vec![<SC::F as Field>::Packing::default(); mc[0].width]);
359        }
360
361        // Pre-allocate DAG node buffer (reused across all packed evaluations)
362        let mut node_buf: Vec<<SC::F as Field>::Packing> =
363            Vec::with_capacity(helper.constraints_dag.nodes.len());
364
365        // Results: d * skip_sz evaluations in EF
366        let mut results: Vec<SC::EF> = vec![SC::EF::ZERO; d * skip_sz];
367
368        // Process z-points in packs of WIDTH
369        for z_base in (0..skip_sz).step_by(width) {
370            let z_count = width.min(skip_sz - z_base);
371
372            for (ci, &zerofier_inv) in zerofier_invs.iter().enumerate() {
373                // Pack WIDTH z-points' column values into F::Packing vectors
374                // Selectors (3 columns)
375                for col in 0..sels_w {
376                    packed_row_parts[0][col] = <SC::F as Field>::Packing::from_fn(|lane| {
377                        if lane < z_count {
378                            sels_cosets[ci].values[(z_base + lane) * sels_w + col]
379                        } else {
380                            SC::F::ZERO
381                        }
382                    });
383                }
384
385                // Trace matrices
386                for (mat_idx, mc) in mat_cosets.iter().enumerate() {
387                    let mw = mc[ci].width;
388                    for col in 0..mw {
389                        packed_row_parts[1 + mat_idx][col] =
390                            <SC::F as Field>::Packing::from_fn(|lane| {
391                                if lane < z_count {
392                                    mc[ci].values[(z_base + lane) * mw + col]
393                                } else {
394                                    SC::F::ZERO
395                                }
396                            });
397                    }
398                }
399
400                // Evaluate DAG once for WIDTH z-points simultaneously
401                let evaluator = helper.evaluator_packed(&packed_row_parts);
402                eval_nodes_into(&evaluator, &helper.constraints_dag.nodes, &mut node_buf);
403
404                // Unpack lanes: accumulate constraint_eval per lane, multiply by eq * zerofier_inv
405                for lane in 0..z_count {
406                    let constraint_eval: SC::EF =
407                        zip(lambda_pows, &helper.constraints_dag.constraint_idx)
408                            .fold(SC::EF::ZERO, |acc, (&lp, &idx)| {
409                                acc + lp * node_buf[idx].as_slice()[lane]
410                            });
411                    let z_idx = z_base + lane;
412                    results[z_idx * d + ci] += eq * constraint_eval * zerofier_inv;
413                }
414            }
415        }
416        results
417    });
418
419    // Reduce: sum over H_n
420    let hypercube_sum = |mut acc: Vec<SC::EF>, x: Vec<SC::EF>| {
421        for (a, b) in acc.iter_mut().zip(x) {
422            *a += b;
423        }
424        acc
425    };
426    cfg_if::cfg_if! {
427        if #[cfg(feature = "parallel")] {
428            let evals = evals.reduce(
429                || vec![SC::EF::ZERO; d << l_skip],
430                hypercube_sum,
431            );
432        } else {
433            let evals: Vec<_> = evals.collect();
434            let evals = evals.into_iter().fold(
435                vec![SC::EF::ZERO; d << l_skip],
436                hypercube_sum,
437            );
438        }
439    }
440
441    // Assemble polynomial from coset evaluations
442    let values: Vec<SC::EF> = evals;
443    UnivariatePoly::from_geometric_cosets_evals_idft(RowMajorMatrix::new(values, d), g, g)
444}
445
446// ============================================================================
447// Constraint evaluator (duplicated from stark-backend since it's pub(super))
448// ============================================================================
449
450struct ViewPair<T> {
451    local: *const T,
452    next: Option<*const T>,
453}
454
455// SAFETY: ViewPair is safe to send between threads as long as the underlying
456// data (pointed to by the raw pointers) is valid and shared immutably.
457unsafe impl<T: Send> Send for ViewPair<T> {}
458unsafe impl<T: Sync> Sync for ViewPair<T> {}
459
460impl<T> ViewPair<T> {
461    fn new(local: &[T], next: Option<&[T]>) -> Self {
462        Self {
463            local: local.as_ptr(),
464            next: next.map(|nxt| nxt.as_ptr()),
465        }
466    }
467
468    /// SAFETY: no matrix bounds checks are done.
469    unsafe fn get(&self, row_offset: usize, column_idx: usize) -> &T {
470        match row_offset {
471            0 => &*self.local.add(column_idx),
472            1 => &*self.next.unwrap_unchecked().add(column_idx),
473            _ => panic!("row offset {row_offset} not supported"),
474        }
475    }
476}
477
478struct ConstraintEvaluator<'a, F, EF> {
479    preprocessed: Option<ViewPair<EF>>,
480    partitioned_main: Vec<ViewPair<EF>>,
481    is_first_row: EF,
482    is_last_row: EF,
483    is_transition: EF,
484    public_values: &'a [F],
485}
486
487impl<F: Field, EF: ExtensionField<F>> SymbolicEvaluator<F, EF> for ConstraintEvaluator<'_, F, EF> {
488    fn eval_const(&self, c: F) -> EF {
489        c.into()
490    }
491    fn eval_is_first_row(&self) -> EF {
492        self.is_first_row
493    }
494    fn eval_is_last_row(&self) -> EF {
495        self.is_last_row
496    }
497    fn eval_is_transition(&self) -> EF {
498        self.is_transition
499    }
500
501    fn eval_var(&self, symbolic_var: SymbolicVariable<F>) -> EF {
502        let index = symbolic_var.index;
503        match symbolic_var.entry {
504            Entry::Preprocessed { offset } => unsafe {
505                *self
506                    .preprocessed
507                    .as_ref()
508                    .unwrap_unchecked()
509                    .get(offset, index)
510            },
511            Entry::Main { part_index, offset } => unsafe {
512                *self.partitioned_main[part_index].get(offset, index)
513            },
514            Entry::Public => unsafe { EF::from(*self.public_values.get_unchecked(index)) },
515            Entry::Challenge => unreachable!("challenge not supported"),
516        }
517    }
518}
519
520// ============================================================================
521// Packed constraint evaluator (SIMD: evaluates DAG in F::Packing)
522// ============================================================================
523
524/// Evaluates the constraint DAG using packed base field (F::Packing), processing
525/// WIDTH independent evaluations simultaneously via SIMD.
526///
527/// On aarch64/NEON: WIDTH=4, on x86/AVX2: WIDTH=8, scalar fallback: WIDTH=1.
528/// This reduces the number of DAG walks from N to ceil(N/WIDTH), cutting
529/// cache misses on the ~30K-node DAG by WIDTH×.
530struct PackedConstraintEvaluator<'a, F: Field> {
531    preprocessed: Option<ViewPair<F::Packing>>,
532    partitioned_main: Vec<ViewPair<F::Packing>>,
533    is_first_row: F::Packing,
534    is_last_row: F::Packing,
535    is_transition: F::Packing,
536    public_values: &'a [F],
537}
538
539impl<F: Field> SymbolicEvaluator<F, F::Packing> for PackedConstraintEvaluator<'_, F> {
540    fn eval_const(&self, c: F) -> F::Packing {
541        F::Packing::from_fn(|_| c)
542    }
543    fn eval_is_first_row(&self) -> F::Packing {
544        self.is_first_row
545    }
546    fn eval_is_last_row(&self) -> F::Packing {
547        self.is_last_row
548    }
549    fn eval_is_transition(&self) -> F::Packing {
550        self.is_transition
551    }
552
553    fn eval_var(&self, symbolic_var: SymbolicVariable<F>) -> F::Packing {
554        let index = symbolic_var.index;
555        match symbolic_var.entry {
556            Entry::Preprocessed { offset } => unsafe {
557                *self
558                    .preprocessed
559                    .as_ref()
560                    .unwrap_unchecked()
561                    .get(offset, index)
562            },
563            Entry::Main { part_index, offset } => unsafe {
564                *self.partitioned_main[part_index].get(offset, index)
565            },
566            Entry::Public => unsafe {
567                F::Packing::from_fn(|_| *self.public_values.get_unchecked(index))
568            },
569            Entry::Challenge => unreachable!("challenge not supported"),
570        }
571    }
572}
573
574/// Evaluate DAG nodes into a pre-allocated buffer, avoiding per-call allocation.
575/// The buffer is cleared and reused across calls, saving ~120KB allocation per
576/// invocation for typical KeccakAir DAGs (~30K nodes).
577#[inline]
578fn eval_nodes_into<F, E>(
579    evaluator: &impl SymbolicEvaluator<F, E>,
580    nodes: &[SymbolicExpressionNode<F>],
581    buf: &mut Vec<E>,
582) where
583    F: Field,
584    E: Add<E, Output = E> + Sub<E, Output = E> + Mul<E, Output = E> + Neg<Output = E> + Clone,
585{
586    buf.clear();
587    for node in nodes {
588        let val = match *node {
589            SymbolicExpressionNode::Variable(var) => evaluator.eval_var(var),
590            SymbolicExpressionNode::Constant(c) => evaluator.eval_const(c),
591            SymbolicExpressionNode::Add {
592                left_idx,
593                right_idx,
594                ..
595            } => buf[left_idx].clone() + buf[right_idx].clone(),
596            SymbolicExpressionNode::Sub {
597                left_idx,
598                right_idx,
599                ..
600            } => buf[left_idx].clone() - buf[right_idx].clone(),
601            SymbolicExpressionNode::Neg { idx, .. } => -buf[idx].clone(),
602            SymbolicExpressionNode::Mul {
603                left_idx,
604                right_idx,
605                ..
606            } => buf[left_idx].clone() * buf[right_idx].clone(),
607            SymbolicExpressionNode::IsFirstRow => evaluator.eval_is_first_row(),
608            SymbolicExpressionNode::IsLastRow => evaluator.eval_is_last_row(),
609            SymbolicExpressionNode::IsTransition => evaluator.eval_is_transition(),
610        };
611        buf.push(val);
612    }
613}
614
615// ============================================================================
616// RowMajorEvalHelper
617// ============================================================================
618
619/// Evaluation helper for a single AIR, storing preprocessed trace as row-major.
620pub(crate) struct RowMajorEvalHelper<'a, SC: StarkProtocolConfig> {
621    pub constraints_dag: &'a SymbolicExpressionDag<SC::F>,
622    pub interactions: Vec<SymbolicInteraction<SC::F>>,
623    pub public_values: Vec<SC::F>,
624    pub preprocessed_trace: Option<&'a RowMajorMatrix<SC::F>>,
625    pub needs_next: bool,
626    pub constraint_degree: u8,
627}
628
629impl<'a, SC: StarkProtocolConfig> RowMajorEvalHelper<'a, SC>
630where
631    SC::F: TwoAdicField,
632    SC::EF: TwoAdicField + ExtensionField<SC::F>,
633{
634    pub fn has_preprocessed(&self) -> bool {
635        self.preprocessed_trace.is_some()
636    }
637
638    /// Returns list of (&RowMajorMatrix, is_rot) in order:
639    /// - (if preprocessed) (preprocessed, false), (preprocessed, true)
640    /// - for each cached: (cached_i, false), (cached_i, true)
641    /// - (common, false), (common, true)
642    pub fn view_mats_rowmaj(
643        &self,
644        ctx: &'a AirProvingContext<CpuBackend<SC>>,
645    ) -> Vec<(&'a RowMajorMatrix<SC::F>, bool)> {
646        let base_mats = usize::from(self.has_preprocessed()) + 1 + ctx.cached_mains.len();
647        let cap = if self.needs_next {
648            2 * base_mats
649        } else {
650            base_mats
651        };
652        let mut mats = Vec::with_capacity(cap);
653        if let Some(pp) = self.preprocessed_trace {
654            mats.push((pp, false));
655            if self.needs_next {
656                mats.push((pp, true));
657            }
658        }
659        for cd in &ctx.cached_mains {
660            mats.push((&cd.trace, false));
661            if self.needs_next {
662                mats.push((&cd.trace, true));
663            }
664        }
665        mats.push((&ctx.common_main, false));
666        if self.needs_next {
667            mats.push((&ctx.common_main, true));
668        }
669        mats
670    }
671
672    /// Build view pairs from row_parts, splitting preprocessed from main partitions.
673    fn build_view_pairs<T>(&self, row_parts: &[Vec<T>]) -> (Option<ViewPair<T>>, Vec<ViewPair<T>>) {
674        let mut view_pairs = if self.needs_next {
675            let mut chunks = row_parts[1..].chunks_exact(2);
676            let pairs = chunks
677                .by_ref()
678                .map(|pair| ViewPair::new(&pair[0], Some(&pair[1][..])))
679                .collect_vec();
680            debug_assert!(chunks.remainder().is_empty());
681            pairs
682        } else {
683            row_parts[1..]
684                .iter()
685                .map(|part| ViewPair::new(part, None))
686                .collect_vec()
687        };
688        let preprocessed = if self.has_preprocessed() {
689            Some(view_pairs.remove(0))
690        } else {
691            None
692        };
693        (preprocessed, view_pairs)
694    }
695
696    fn evaluator<FF: ExtensionField<SC::F>>(
697        &self,
698        row_parts: &[Vec<FF>],
699    ) -> ConstraintEvaluator<'_, SC::F, FF> {
700        let sels = &row_parts[0];
701        let (preprocessed, partitioned_main) = self.build_view_pairs(row_parts);
702        ConstraintEvaluator {
703            preprocessed,
704            partitioned_main,
705            is_first_row: sels[0],
706            is_transition: sels[1],
707            is_last_row: sels[2],
708            public_values: &self.public_values,
709        }
710    }
711
712    pub fn acc_constraints<FF: ExtensionField<SC::F>, EF: ExtensionField<FF>>(
713        &self,
714        row_parts: &[Vec<FF>],
715        lambda_pows: &[EF],
716    ) -> EF {
717        let evaluator = self.evaluator(row_parts);
718        let nodes = evaluator.eval_nodes(&self.constraints_dag.nodes);
719        zip(lambda_pows, &self.constraints_dag.constraint_idx)
720            .fold(EF::ZERO, |acc, (&lambda_pow, &idx)| {
721                acc + lambda_pow * nodes[idx]
722            })
723    }
724
725    fn evaluator_packed(
726        &self,
727        row_parts: &[Vec<<SC::F as Field>::Packing>],
728    ) -> PackedConstraintEvaluator<'_, SC::F> {
729        let sels = &row_parts[0];
730        let (preprocessed, partitioned_main) = self.build_view_pairs(row_parts);
731        PackedConstraintEvaluator {
732            preprocessed,
733            partitioned_main,
734            is_first_row: sels[0],
735            is_transition: sels[1],
736            is_last_row: sels[2],
737            public_values: &self.public_values,
738        }
739    }
740
741    pub fn acc_interactions<FF, EF>(
742        &self,
743        row_parts: &[Vec<FF>],
744        beta_pows: &[EF],
745        eq_3bs: &[EF],
746    ) -> [EF; 2]
747    where
748        FF: ExtensionField<SC::F>,
749        EF: ExtensionField<FF> + ExtensionField<SC::F>,
750    {
751        let interaction_evals = self.eval_interactions(row_parts, beta_pows);
752        let mut numer = EF::ZERO;
753        let mut denom = EF::ZERO;
754        for (&eq_3b, eval) in zip(eq_3bs, interaction_evals) {
755            numer += eq_3b * eval.0;
756            denom += eq_3b * eval.1;
757        }
758        [numer, denom]
759    }
760
761    pub fn eval_interactions<FF, EF>(
762        &self,
763        row_parts: &[Vec<FF>],
764        beta_pows: &[EF],
765    ) -> Vec<(FF, EF)>
766    where
767        FF: ExtensionField<SC::F>,
768        EF: ExtensionField<FF> + ExtensionField<SC::F>,
769    {
770        let evaluator = self.evaluator(row_parts);
771        self.interactions
772            .iter()
773            .map(|interaction| {
774                let b = SC::F::from_u32(interaction.bus_index as u32 + 1);
775                let msg_len = interaction.message.len();
776                assert!(msg_len <= beta_pows.len());
777                let denom = zip(&interaction.message, beta_pows).fold(
778                    beta_pows[msg_len] * b,
779                    |h_beta, (msg_j, &beta_j)| {
780                        let msg_j_eval = evaluator.eval_expr(msg_j);
781                        h_beta + beta_j * msg_j_eval
782                    },
783                );
784                let numer = evaluator.eval_expr(&interaction.count);
785                (numer, denom)
786            })
787            .collect()
788    }
789
790    /// Build row_parts from row-major matrices for a given row index.
791    /// This is the performance-critical row-major access pattern:
792    /// each row is contiguous in memory.
793    fn build_row_parts(
794        mats: &[(&RowMajorMatrix<SC::F>, bool)],
795        row_idx: usize,
796        height: usize,
797    ) -> Vec<Vec<SC::F>> {
798        let is_first = SC::F::from_bool(row_idx == 0);
799        let is_transition = SC::F::from_bool(row_idx != height - 1);
800        let is_last = SC::F::from_bool(row_idx == height - 1);
801
802        let mut row_parts = Vec::with_capacity(mats.len() + 1);
803        row_parts.push(vec![is_first, is_transition, is_last]);
804
805        for &(mat, is_rot) in mats {
806            let mat_height = mat.values.len() / mat.width;
807            let idx = if is_rot {
808                (row_idx + 1) % mat_height
809            } else {
810                row_idx % mat_height
811            };
812            let start = idx * mat.width;
813            // Contiguous row access — the key cache locality optimization
814            row_parts.push(mat.values[start..start + mat.width].to_vec());
815        }
816        row_parts
817    }
818}
819
820// ============================================================================
821// LogupZerocheckRowMajor
822// ============================================================================
823
824pub(crate) struct LogupZerocheckRowMajor<'a, SC: StarkProtocolConfig> {
825    pub beta_pows: Vec<SC::EF>,
826
827    pub l_skip: usize,
828    pub n_logup: usize,
829
830    pub omega_skip_pows: Vec<SC::F>,
831
832    pub interactions_layout: StackedLayout,
833    pub(crate) eval_helpers: Vec<RowMajorEvalHelper<'a, SC>>,
834    pub constraint_degree: usize,
835    pub n_per_trace: Vec<isize>,
836    max_num_constraints: usize,
837
838    pub xi: Vec<SC::EF>,
839    lambda_pows: Vec<SC::EF>,
840    eq_xi_per_trace: Vec<Vec<SC::EF>>,
841    eq_3b_per_trace: Vec<Vec<SC::EF>>,
842    sels_per_trace_base: Vec<RowMajorMatrix<SC::F>>,
843    pub mat_evals_per_trace: Vec<Vec<RowMajorMatrix<SC::EF>>>,
844    pub sels_per_trace: Vec<RowMajorMatrix<SC::EF>>,
845    pub(crate) zerocheck_tilde_evals: Vec<SC::EF>,
846    pub(crate) logup_tilde_evals: Vec<[SC::EF; 2]>,
847
848    pub(crate) prev_s_eval: SC::EF,
849    pub(crate) eq_ns: Vec<SC::EF>,
850    pub(crate) eq_sharp_ns: Vec<SC::EF>,
851}
852
853impl<'a, SC: StarkProtocolConfig> LogupZerocheckRowMajor<'a, SC>
854where
855    SC::F: TwoAdicField,
856    SC::EF: TwoAdicField + ExtensionField<SC::F>,
857    CpuBackend<SC>: ProverBackend<Val = SC::F, Matrix = RowMajorMatrix<SC::F>>,
858{
859    pub fn new(
860        pk: &'a DeviceMultiStarkProvingKey<CpuBackend<SC>>,
861        ctx: &ProvingContext<CpuBackend<SC>>,
862        n_logup: usize,
863        interactions_layout: StackedLayout,
864        _alpha_logup: SC::EF,
865        beta_logup: SC::EF,
866    ) -> Self {
867        let l_skip = pk.params.l_skip;
868        let omega_skip = SC::F::two_adic_generator(l_skip);
869        let omega_skip_pows = omega_skip.powers().take(1 << l_skip).collect_vec();
870        let num_airs_present = ctx.per_trace.len();
871
872        let constraint_degree = pk.max_constraint_degree;
873        let max_interaction_length = ctx
874            .per_trace
875            .iter()
876            .flat_map(|(air_idx, _)| {
877                pk.per_air[*air_idx]
878                    .vk
879                    .symbolic_constraints
880                    .interactions
881                    .iter()
882                    .map(|i| i.message.len())
883            })
884            .max()
885            .unwrap_or(0);
886        let beta_pows = beta_logup
887            .powers()
888            .take(max_interaction_length + 1)
889            .collect_vec();
890
891        let n_per_trace: Vec<isize> = ctx
892            .common_main_traces()
893            .map(|(_, t)| log2_strict_usize(MatrixDimensions::height(t)) as isize - l_skip as isize)
894            .collect_vec();
895        let n_max: usize = n_per_trace[0].max(0) as usize;
896
897        let eval_helpers: Vec<RowMajorEvalHelper<SC>> = ctx
898            .per_trace
899            .iter()
900            .map(|(air_idx, trace_ctx)| {
901                let pk = &pk.per_air[*air_idx];
902                let constraints = &pk.vk.symbolic_constraints.constraints;
903                let public_values = trace_ctx.public_values.clone();
904                let preprocessed_trace: Option<&RowMajorMatrix<SC::F>> =
905                    pk.preprocessed_data.as_ref().map(|cd| &cd.trace);
906                // Validate index bounds
907                let mut rotation = 0;
908                for node in &constraints.nodes {
909                    if let SymbolicExpressionNode::Variable(var) = node {
910                        match var.entry {
911                            Entry::Preprocessed { offset } => {
912                                rotation = max(rotation, offset);
913                                assert!(
914                                    var.index < preprocessed_trace.unwrap().width,
915                                    "col_index={} >= preprocessed width={}",
916                                    var.index,
917                                    preprocessed_trace.unwrap().width
918                                );
919                            }
920                            Entry::Main { part_index, offset } => {
921                                rotation = max(rotation, offset);
922                                // Get width of the partition
923                                let part_width = if part_index < trace_ctx.cached_mains.len() {
924                                    trace_ctx.cached_mains[part_index].trace.width
925                                } else {
926                                    trace_ctx.common_main.width
927                                };
928                                assert!(
929                                    var.index < part_width,
930                                    "col_index={} >= main partition {} width={}",
931                                    var.index,
932                                    part_index,
933                                    part_width
934                                );
935                            }
936                            Entry::Public => {
937                                assert!(var.index < public_values.len());
938                            }
939                            Entry::Challenge => unreachable!("challenge not supported"),
940                        }
941                    }
942                }
943                let needs_next = pk.vk.params.need_rot;
944                debug_assert_eq!(needs_next, rotation > 0);
945                let symbolic_constraints = SymbolicConstraints::from(&pk.vk.symbolic_constraints);
946                RowMajorEvalHelper {
947                    constraints_dag: &pk.vk.symbolic_constraints.constraints,
948                    interactions: symbolic_constraints.interactions,
949                    public_values,
950                    preprocessed_trace,
951                    needs_next,
952                    constraint_degree: pk.vk.max_constraint_degree,
953                }
954            })
955            .collect();
956
957        let max_num_constraints = pk
958            .per_air
959            .iter()
960            .map(|pk| pk.vk.symbolic_constraints.constraints.constraint_idx.len())
961            .max()
962            .unwrap_or(0);
963
964        let zerocheck_tilde_evals = vec![SC::EF::ZERO; num_airs_present];
965        let logup_tilde_evals = vec![[SC::EF::ZERO; 2]; num_airs_present];
966        Self {
967            beta_pows,
968            l_skip,
969            n_logup,
970            omega_skip_pows,
971            interactions_layout,
972            constraint_degree,
973            max_num_constraints,
974            n_per_trace,
975            eval_helpers,
976            xi: vec![],
977            lambda_pows: vec![],
978            sels_per_trace_base: vec![],
979            eq_xi_per_trace: vec![],
980            eq_3b_per_trace: vec![],
981            mat_evals_per_trace: vec![],
982            sels_per_trace: vec![],
983            zerocheck_tilde_evals,
984            logup_tilde_evals,
985            prev_s_eval: SC::EF::ZERO,
986            eq_ns: Vec::with_capacity(n_max + 1),
987            eq_sharp_ns: Vec::with_capacity(n_max + 1),
988        }
989    }
990
991    pub fn sumcheck_uni_round0_polys(
992        &mut self,
993        ctx: &ProvingContext<CpuBackend<SC>>,
994        lambda: SC::EF,
995    ) -> Vec<UnivariatePoly<SC::EF>> {
996        let n_logup = self.n_logup;
997        let l_skip = self.l_skip;
998        let xi = &self.xi;
999        self.lambda_pows = lambda.powers().take(self.max_num_constraints).collect_vec();
1000
1001        self.eq_3b_per_trace = self
1002            .eval_helpers
1003            .par_iter()
1004            .zip(&self.n_per_trace)
1005            .enumerate()
1006            .map(|(trace_idx, (helper, &n))| {
1007                let n_lift = n.max(0) as usize;
1008                if helper.interactions.is_empty() {
1009                    return vec![];
1010                }
1011                let mut b_vec = vec![SC::F::ZERO; n_logup - n_lift];
1012                (0..helper.interactions.len())
1013                    .map(|i| {
1014                        let stacked_idx =
1015                            self.interactions_layout.get(trace_idx, i).unwrap().row_idx;
1016                        debug_assert!(stacked_idx.trailing_zeros() as usize >= n_lift + l_skip);
1017                        let mut b_int = stacked_idx >> (l_skip + n_lift);
1018                        for b in &mut b_vec {
1019                            *b = SC::F::from_bool(b_int & 1 == 1);
1020                            b_int >>= 1;
1021                        }
1022                        eval_eq_mle(&xi[l_skip + n_lift..l_skip + n_logup], &b_vec)
1023                    })
1024                    .collect_vec()
1025            })
1026            .collect::<Vec<_>>();
1027
1028        self.eq_xi_per_trace = self
1029            .n_per_trace
1030            .par_iter()
1031            .map(|&n| {
1032                let n_lift = n.max(0) as usize;
1033                evals_eq_hypercubes(n_lift, xi[l_skip..l_skip + n_lift].iter().rev())
1034            })
1035            .collect();
1036
1037        self.sels_per_trace_base = self
1038            .n_per_trace
1039            .iter()
1040            .map(|&n| {
1041                let log_height = l_skip.checked_add_signed(n).unwrap();
1042                let height = 1 << log_height;
1043                let lifted_height = height.max(1 << l_skip);
1044                // Row-major: each row is [is_first, is_transition, is_last]
1045                let mut vals = SC::F::zero_vec(3 * lifted_height);
1046                for i in 0..lifted_height {
1047                    let row_in_period = i % height;
1048                    vals[i * 3] = SC::F::from_bool(row_in_period == 0); // is_first
1049                    vals[i * 3 + 1] = SC::F::from_bool(row_in_period != height - 1); // is_transition
1050                    vals[i * 3 + 2] = SC::F::from_bool(row_in_period == height - 1);
1051                    // is_last
1052                }
1053                RowMajorMatrix::new(vals, 3)
1054            })
1055            .collect_vec();
1056
1057        // Zerocheck round 0 — uses batch DFT from row-major data
1058        let sp_0_zerochecks = self
1059            .eval_helpers
1060            .par_iter()
1061            .enumerate()
1062            .map(|(trace_idx, helper)| {
1063                let trace_ctx = &ctx.per_trace[trace_idx].1;
1064                let n_lift = log2_strict_usize(trace_ctx.height()).saturating_sub(l_skip);
1065                let rm_mats = helper.view_mats_rowmaj(trace_ctx);
1066                let eq_xi = &self.eq_xi_per_trace[trace_idx][(1 << n_lift) - 1..(2 << n_lift) - 1];
1067                let sels_cm = &self.sels_per_trace_base[trace_idx];
1068
1069                let constraint_deg = helper.constraint_degree as usize;
1070                if constraint_deg == 0 {
1071                    return UnivariatePoly::new(vec![]);
1072                }
1073                let num_cosets = constraint_deg - 1;
1074                let q = sumcheck_uni_round0_zerocheck_packed::<SC>(
1075                    l_skip,
1076                    n_lift,
1077                    num_cosets,
1078                    &rm_mats,
1079                    sels_cm,
1080                    helper,
1081                    eq_xi,
1082                    &self.lambda_pows,
1083                );
1084                let sp_0_deg = sumcheck_round0_deg(l_skip, constraint_deg);
1085                let coeffs = (0..=sp_0_deg)
1086                    .map(|i| {
1087                        let mut c = -*q.coeffs().get(i).unwrap_or(&SC::EF::ZERO);
1088                        if i >= 1 << l_skip {
1089                            c += q.coeffs()[i - (1 << l_skip)];
1090                        }
1091                        c
1092                    })
1093                    .collect_vec();
1094                debug_assert_eq!(
1095                    coeffs.iter().step_by(1 << l_skip).copied().sum::<SC::EF>(),
1096                    SC::EF::ZERO,
1097                    "Zerocheck sum is not zero for air_id: {}",
1098                    ctx.per_trace[trace_idx].0
1099                );
1100                UnivariatePoly::new(coeffs)
1101            })
1102            .collect::<Vec<_>>();
1103
1104        // Logup round 0 — uses batch DFT from row-major data
1105        let sp_0_logups = self
1106            .eval_helpers
1107            .par_iter()
1108            .enumerate()
1109            .flat_map(|(trace_idx, helper)| {
1110                if helper.interactions.is_empty() {
1111                    return [(); 2].map(|_| UnivariatePoly::new(vec![]));
1112                }
1113                let trace_ctx = &ctx.per_trace[trace_idx].1;
1114                let log_height = log2_strict_usize(trace_ctx.height());
1115                let n_lift = log_height.saturating_sub(l_skip);
1116                let rm_mats = helper.view_mats_rowmaj(trace_ctx);
1117                let eq_xi = &self.eq_xi_per_trace[trace_idx][(1 << n_lift) - 1..(2 << n_lift) - 1];
1118                let eq_3bs = &self.eq_3b_per_trace[trace_idx];
1119                let sels_cm = &self.sels_per_trace_base[trace_idx];
1120                let norm_factor_denom = 1 << l_skip.saturating_sub(log_height);
1121                let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
1122
1123                let [mut numer, denom] = sumcheck_uni_round0_batch::<SC::F, SC::EF, _, 2>(
1124                    l_skip,
1125                    n_lift,
1126                    helper.constraint_degree as usize,
1127                    &rm_mats,
1128                    sels_cm,
1129                    |_z, x, row_parts| {
1130                        let eq = eq_xi[x];
1131                        let [numer, denom] =
1132                            helper.acc_interactions(row_parts, &self.beta_pows, eq_3bs);
1133                        [eq * numer, eq * denom]
1134                    },
1135                );
1136                for p in numer.coeffs_mut() {
1137                    *p *= norm_factor;
1138                }
1139                [numer, denom]
1140            })
1141            .collect::<Vec<_>>();
1142
1143        sp_0_logups.into_iter().chain(sp_0_zerochecks).collect()
1144    }
1145
1146    /// Fold PLE evaluations using randomness r_0.
1147    /// Reads directly from row-major matrices via barycentric interpolation,
1148    /// avoiding the O(n*m) transpose to col-major (saves ~3s for keccakf).
1149    pub fn fold_ple_evals(&mut self, ctx: &ProvingContext<CpuBackend<SC>>, r_0: SC::EF) {
1150        let l_skip = self.l_skip;
1151        self.mat_evals_per_trace = self
1152            .eval_helpers
1153            .par_iter()
1154            .zip(ctx.per_trace.par_iter())
1155            .map(|(helper, (_, trace_ctx))| {
1156                let rm_mats = helper.view_mats_rowmaj(trace_ctx);
1157                rm_mats
1158                    .into_iter()
1159                    .map(|(rm, is_rot)| fold_ple_evals_rowmajor(l_skip, rm, is_rot, r_0))
1160                    .collect::<Vec<_>>()
1161            })
1162            .collect::<Vec<_>>();
1163        self.sels_per_trace = take(&mut self.sels_per_trace_base)
1164            .iter()
1165            .map(|rm| fold_ple_evals_rowmajor(l_skip, rm, false, r_0))
1166            .collect();
1167        let eq_r0 = eval_eq_uni(l_skip, self.xi[0], r_0);
1168        let eq_sharp_r0 = eval_eq_sharp_uni(&self.omega_skip_pows, &self.xi[..l_skip], r_0);
1169        self.eq_ns.push(eq_r0);
1170        self.eq_sharp_ns.push(eq_sharp_r0);
1171        self.eq_xi_per_trace.iter_mut().for_each(|eq| {
1172            if eq.len() > 1 {
1173                eq.truncate(eq.len() / 2);
1174            }
1175        });
1176    }
1177
1178    /// After PLE folding, operates on small RowMajorMatrix<EF> via RM sumcheck.
1179    pub fn sumcheck_polys_eval(&mut self, round: usize, r_prev: SC::EF) -> Vec<Vec<SC::EF>> {
1180        let sp_deg = self.constraint_degree;
1181        let sp_zerocheck_evals: Vec<Vec<SC::EF>> = izip!(
1182            &self.eval_helpers,
1183            &mut self.zerocheck_tilde_evals,
1184            &self.n_per_trace,
1185            &self.mat_evals_per_trace,
1186            &self.sels_per_trace,
1187            &self.eq_xi_per_trace
1188        )
1189        .map(|(helper, tilde_eval, &n, mats, sels, eq_xi_tree)| {
1190            let n_lift = n.max(0) as usize;
1191            if round > n_lift {
1192                if round == n_lift + 1 {
1193                    // Height is 1 — extract single row as Vec
1194                    let parts: Vec<Vec<SC::EF>> = iter::once(sels)
1195                        .chain(mats.iter())
1196                        .map(|mat| mat.values.to_vec())
1197                        .collect();
1198                    let eq_r_acc = *self.eq_ns.last().unwrap();
1199                    *tilde_eval = eq_r_acc * helper.acc_constraints(&parts, &self.lambda_pows);
1200                } else {
1201                    *tilde_eval *= r_prev;
1202                };
1203                vec![*tilde_eval]
1204            } else {
1205                let log_num_y = n_lift - round;
1206                let num_y = 1 << log_num_y;
1207                let eq_xi = &eq_xi_tree[num_y - 1..];
1208                let parts_vec: Vec<&RowMajorMatrix<SC::EF>> =
1209                    iter::once(sels).chain(mats.iter()).collect();
1210                let [s] = crate::row_major_ops::sumcheck_round_poly_evals_rm(
1211                    log_num_y + 1,
1212                    sp_deg,
1213                    &parts_vec,
1214                    |_x, y, row_parts| {
1215                        let eq = eq_xi[y];
1216                        let constraint_eval = helper.acc_constraints(row_parts, &self.lambda_pows);
1217                        [eq * constraint_eval]
1218                    },
1219                );
1220                s
1221            }
1222        })
1223        .collect();
1224
1225        let sp_logup_evals: Vec<Vec<SC::EF>> = izip!(
1226            &self.eval_helpers,
1227            &mut self.logup_tilde_evals,
1228            &self.n_per_trace,
1229            &self.mat_evals_per_trace,
1230            &self.sels_per_trace,
1231            &self.eq_xi_per_trace,
1232            &self.eq_3b_per_trace
1233        )
1234        .flat_map(|(helper, tilde_eval, &n, mats, sels, eq_xi_tree, eq_3bs)| {
1235            if helper.interactions.is_empty() {
1236                return [vec![SC::EF::ZERO; sp_deg], vec![SC::EF::ZERO; sp_deg]];
1237            }
1238            let n_lift = n.max(0) as usize;
1239            let norm_factor_denom = 1 << (-n).max(0);
1240            let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
1241            if round > n_lift {
1242                if round == n_lift + 1 {
1243                    // Height is 1 — extract single row as Vec
1244                    let parts: Vec<Vec<SC::EF>> = iter::once(sels)
1245                        .chain(mats.iter())
1246                        .map(|mat| mat.values.to_vec())
1247                        .collect();
1248                    let eq_sharp_r_acc = *self.eq_sharp_ns.last().unwrap();
1249                    *tilde_eval = helper
1250                        .acc_interactions(&parts, &self.beta_pows, eq_3bs)
1251                        .map(|x| eq_sharp_r_acc * x);
1252                    tilde_eval[0] *= norm_factor;
1253                } else {
1254                    for x in tilde_eval.iter_mut() {
1255                        *x *= r_prev;
1256                    }
1257                };
1258                tilde_eval.map(|tilde_eval| vec![tilde_eval])
1259            } else {
1260                let parts_vec: Vec<&RowMajorMatrix<SC::EF>> =
1261                    iter::once(sels).chain(mats.iter()).collect();
1262                let log_num_y = n_lift - round;
1263                let num_y = 1 << log_num_y;
1264                let eq_xi = &eq_xi_tree[num_y - 1..];
1265                let [mut numer, denom] = crate::row_major_ops::sumcheck_round_poly_evals_rm(
1266                    log_num_y + 1,
1267                    sp_deg,
1268                    &parts_vec,
1269                    |_x, y, row_parts| {
1270                        let eq = eq_xi[y];
1271                        helper
1272                            .acc_interactions(row_parts, &self.beta_pows, eq_3bs)
1273                            .map(|eval| eq * eval)
1274                    },
1275                );
1276                for p in &mut numer {
1277                    *p *= norm_factor;
1278                }
1279                [numer, denom]
1280            }
1281        })
1282        .collect();
1283
1284        sp_logup_evals
1285            .into_iter()
1286            .chain(sp_zerocheck_evals)
1287            .collect()
1288    }
1289
1290    pub fn fold_mle_evals(&mut self, round: usize, r_round: SC::EF) {
1291        self.mat_evals_per_trace = take(&mut self.mat_evals_per_trace)
1292            .into_iter()
1293            .map(|mats| crate::row_major_ops::batch_fold_mle_evals_rm(mats, r_round))
1294            .collect_vec();
1295        self.sels_per_trace =
1296            crate::row_major_ops::batch_fold_mle_evals_rm(take(&mut self.sels_per_trace), r_round);
1297        self.eq_xi_per_trace.par_iter_mut().for_each(|eq| {
1298            if eq.len() > 1 {
1299                eq.truncate(eq.len() / 2);
1300            }
1301        });
1302        let xi = self.xi[self.l_skip + round - 1];
1303        let eq_r = eval_eq_mle(&[xi], &[r_round]);
1304        self.eq_ns.push(self.eq_ns[round - 1] * eq_r);
1305        self.eq_sharp_ns.push(self.eq_sharp_ns[round - 1] * eq_r);
1306    }
1307
1308    pub fn into_column_openings(&mut self) -> Vec<Vec<Vec<SC::EF>>> {
1309        let num_airs_present = self.mat_evals_per_trace.len();
1310        let mut column_openings = Vec::with_capacity(num_airs_present);
1311        for (helper, mut mat_evals) in self
1312            .eval_helpers
1313            .iter()
1314            .zip(take(&mut self.mat_evals_per_trace))
1315        {
1316            // All matrices have height 1 at this point — values IS the single row
1317            let openings_of_air: Vec<Vec<SC::EF>> = if helper.needs_next {
1318                let common_main_rot = mat_evals.pop().unwrap();
1319                let common_main = mat_evals.pop().unwrap();
1320                iter::once(&[common_main, common_main_rot] as &[_])
1321                    .chain(mat_evals.chunks_exact(2))
1322                    .map(|pair| {
1323                        pair[0]
1324                            .values
1325                            .iter()
1326                            .zip(pair[1].values.iter())
1327                            .flat_map(|(&claim, &claim_rot)| [claim, claim_rot])
1328                            .collect_vec()
1329                    })
1330                    .collect_vec()
1331            } else {
1332                let common_main = mat_evals.pop().unwrap();
1333                iter::once(common_main)
1334                    .chain(mat_evals.into_iter())
1335                    .map(|mat| mat.values)
1336                    .collect_vec()
1337            };
1338            column_openings.push(openings_of_air);
1339        }
1340        column_openings
1341    }
1342}
1343
1344// ============================================================================
1345// prove_zerocheck_and_logup
1346// ============================================================================
1347
1348#[instrument(level = "info", skip_all)]
1349pub fn prove_zerocheck_and_logup<SC: StarkProtocolConfig, TS>(
1350    transcript: &mut TS,
1351    mpk: &DeviceMultiStarkProvingKey<CpuBackend<SC>>,
1352    ctx: &ProvingContext<CpuBackend<SC>>,
1353) -> Result<(GkrProof<SC>, BatchConstraintProof<SC>, Vec<SC::EF>), LogupZerocheckError>
1354where
1355    TS: FiatShamirTranscript<SC>,
1356    SC::F: TwoAdicField,
1357    SC::EF: TwoAdicField + ExtensionField<SC::F>,
1358    CpuBackend<SC>: ProverBackend<Val = SC::F, Matrix = RowMajorMatrix<SC::F>>,
1359{
1360    let l_skip = mpk.params.l_skip;
1361    let constraint_degree = mpk.max_constraint_degree;
1362    let num_traces = ctx.per_trace.len();
1363
1364    let n_max = log2_strict_usize(MatrixDimensions::height(&ctx.per_trace[0].1.common_main))
1365        .saturating_sub(l_skip);
1366    let mut total_interactions = 0u64;
1367    let interactions_meta: Vec<_> = ctx
1368        .per_trace
1369        .iter()
1370        .map(|(air_idx, trace_ctx)| {
1371            let pk = &mpk.per_air[*air_idx];
1372            let num_interactions = pk.vk.symbolic_constraints.interactions.len();
1373            let height = MatrixDimensions::height(&trace_ctx.common_main);
1374            let log_height = log2_strict_usize(height);
1375            let log_lifted_height = log_height.max(l_skip);
1376            total_interactions += (num_interactions as u64) << log_lifted_height;
1377            (num_interactions, log_lifted_height)
1378        })
1379        .collect();
1380    let n_logup = calculate_n_logup(l_skip, total_interactions);
1381    debug!(%n_logup);
1382    let interactions_layout = StackedLayout::new(0, l_skip + n_logup, interactions_meta)?;
1383
1384    let logup_pow_witness = transcript.grind(mpk.params.logup.pow_bits);
1385    let alpha_logup = transcript.sample_ext();
1386    let beta_logup = transcript.sample_ext();
1387    debug!(%alpha_logup, %beta_logup);
1388
1389    let mut prover = LogupZerocheckRowMajor::new(
1390        mpk,
1391        ctx,
1392        n_logup,
1393        interactions_layout,
1394        alpha_logup,
1395        beta_logup,
1396    );
1397
1398    // GKR: compute logup input layer using row-major access
1399    let has_interactions = !prover.interactions_layout.sorted_cols.is_empty();
1400    let gkr_input_evals = if !has_interactions {
1401        vec![]
1402    } else {
1403        // Per trace: row-major interaction evaluations using contiguous row access
1404        let unstacked_interaction_evals = prover
1405            .eval_helpers
1406            .par_iter()
1407            .enumerate()
1408            .map(|(trace_idx, helper)| {
1409                let trace_ctx = &ctx.per_trace[trace_idx].1;
1410                let mats = helper.view_mats_rowmaj(trace_ctx);
1411                let height = MatrixDimensions::height(&trace_ctx.common_main);
1412                (0..height)
1413                    .into_par_iter()
1414                    .map(|i| {
1415                        // Build row_parts from contiguous row-major memory
1416                        let row_parts = RowMajorEvalHelper::<SC>::build_row_parts(&mats, i, height);
1417                        helper.eval_interactions(&row_parts, &prover.beta_pows)
1418                    })
1419                    .collect::<Vec<_>>()
1420            })
1421            .collect::<Vec<_>>();
1422        let mut evals = vec![Frac::default(); 1 << (l_skip + n_logup)];
1423        for (trace_idx, interaction_idx, s) in
1424            prover.interactions_layout.sorted_cols.iter().copied()
1425        {
1426            let pq_evals = &unstacked_interaction_evals[trace_idx];
1427            let height = pq_evals.len();
1428            debug_assert_eq!(s.col_idx, 0);
1429            debug_assert_eq!(1 << s.log_height(), s.len(0));
1430            debug_assert_eq!(s.len(0) % height, 0);
1431            let norm_factor_denom = s.len(0) / height;
1432            let norm_factor = SC::F::from_usize(norm_factor_denom).inverse();
1433            evals[s.row_idx..s.row_idx + s.len(0)]
1434                .chunks_exact_mut(height)
1435                .for_each(|evals| {
1436                    evals
1437                        .par_iter_mut()
1438                        .zip(pq_evals)
1439                        .for_each(|(pq_eval, evals_at_z)| {
1440                            let (mut numer, denom) = evals_at_z[interaction_idx];
1441                            numer *= norm_factor;
1442                            *pq_eval = Frac::new(numer.into(), denom);
1443                        });
1444                });
1445        }
1446        evals.par_iter_mut().for_each(|frac| frac.q += alpha_logup);
1447        evals
1448    };
1449
1450    let (frac_sum_proof, mut xi) =
1451        fractional_sumcheck::<SC, _>(transcript, &gkr_input_evals, true)?;
1452
1453    let n_global = max(n_max, n_logup);
1454    debug!(%n_global);
1455    while xi.len() != l_skip + n_global {
1456        xi.push(transcript.sample_ext());
1457    }
1458    debug!(?xi);
1459    prover.xi = xi;
1460
1461    // Begin batch sumcheck
1462    let mut sumcheck_round_polys = Vec::with_capacity(n_max);
1463    let mut r = Vec::with_capacity(n_max + 1);
1464    let lambda = transcript.sample_ext();
1465    debug!(%lambda);
1466
1467    let sp_0_polys = prover.sumcheck_uni_round0_polys(ctx, lambda);
1468    let sp_0_deg = sumcheck_round0_deg(l_skip, constraint_degree);
1469    let s_deg = constraint_degree + 1;
1470    let s_0_deg = sumcheck_round0_deg(l_skip, s_deg);
1471    let large_uni_domain = (s_0_deg + 1).next_power_of_two();
1472    let dft = Radix2BowersSerial;
1473    let s_0_logup_polys = {
1474        let eq_sharp_uni = eq_sharp_uni_poly(&prover.xi[..l_skip]);
1475        let mut eq_coeffs = eq_sharp_uni.into_coeffs();
1476        eq_coeffs.resize(large_uni_domain, SC::EF::ZERO);
1477        let eq_evals = dft.dft(eq_coeffs);
1478
1479        let width = 2 * num_traces;
1480        let mut sp_coeffs_mat = SC::EF::zero_vec(width * large_uni_domain);
1481        for (i, coeffs) in sp_0_polys[..2 * num_traces].iter().enumerate() {
1482            for (j, &c_j) in coeffs.coeffs().iter().enumerate().take(sp_0_deg + 1) {
1483                unsafe {
1484                    *sp_coeffs_mat.get_unchecked_mut(j * width + i) = c_j;
1485                }
1486            }
1487        }
1488        let mut s_evals = dft.dft_batch(RowMajorMatrix::new(sp_coeffs_mat, width));
1489        for (eq, row) in zip(eq_evals, s_evals.values.chunks_mut(width)) {
1490            for x in row {
1491                *x *= eq;
1492            }
1493        }
1494        dft.idft_batch(s_evals)
1495    };
1496
1497    let skip_domain_size = SC::F::from_usize(1 << l_skip);
1498    let (numerator_term_per_air, denominator_term_per_air): (Vec<_>, Vec<_>) = (0..num_traces)
1499        .map(|trace_idx| {
1500            let [sum_claim_p, sum_claim_q] = [0, 1].map(|is_denom| {
1501                (0..=s_0_deg)
1502                    .step_by(1 << l_skip)
1503                    .map(|j| unsafe {
1504                        *s_0_logup_polys
1505                            .values
1506                            .get_unchecked(j * 2 * num_traces + 2 * trace_idx + is_denom)
1507                    })
1508                    .sum::<SC::EF>()
1509                    * skip_domain_size
1510            });
1511            transcript.observe_ext(sum_claim_p);
1512            transcript.observe_ext(sum_claim_q);
1513            (sum_claim_p, sum_claim_q)
1514        })
1515        .unzip();
1516
1517    let mu = transcript.sample_ext();
1518    debug!(%mu);
1519    let mu_pows = mu.powers().take(3 * num_traces).collect_vec();
1520
1521    let s_0_zc_poly = {
1522        let eq_uni = eq_uni_poly::<SC::F, _>(l_skip, prover.xi[0]);
1523        let mut eq_coeffs = eq_uni.into_coeffs();
1524        eq_coeffs.resize(large_uni_domain, SC::EF::ZERO);
1525        let eq_evals = dft.dft(eq_coeffs);
1526
1527        let mut sp_coeffs = SC::EF::zero_vec(large_uni_domain);
1528        let mus = &mu_pows[2 * num_traces..];
1529        let polys = &sp_0_polys[2 * num_traces..];
1530        for (j, batch_coeff) in sp_coeffs.iter_mut().enumerate().take(sp_0_deg + 1) {
1531            for (&mu, poly) in zip(mus, polys) {
1532                *batch_coeff += mu * *poly.coeffs().get(j).unwrap_or(&SC::EF::ZERO);
1533            }
1534        }
1535        let mut s_evals = dft.dft(sp_coeffs);
1536        for (eq, x) in zip(eq_evals, &mut s_evals) {
1537            *x *= eq;
1538        }
1539        dft.idft(s_evals)
1540    };
1541
1542    let s_0_poly = UnivariatePoly::new(
1543        zip(
1544            s_0_logup_polys.values.chunks_exact(2 * num_traces),
1545            s_0_zc_poly,
1546        )
1547        .take(s_0_deg + 1)
1548        .map(|(logup_row, batched_zc)| {
1549            let coeff = batched_zc
1550                + zip(&mu_pows, logup_row)
1551                    .map(|(&mu_j, &x)| mu_j * x)
1552                    .sum::<SC::EF>();
1553            transcript.observe_ext(coeff);
1554            coeff
1555        })
1556        .collect(),
1557    );
1558
1559    let r_0 = transcript.sample_ext();
1560    r.push(r_0);
1561    debug!(round = 0, r_round = %r_0);
1562    prover.prev_s_eval = s_0_poly.eval_at_point(r_0);
1563    debug!("s_0(r_0) = {}", prover.prev_s_eval);
1564
1565    prover.fold_ple_evals(ctx, r_0);
1566
1567    // MLE rounds
1568    let _mle_rounds_span =
1569        info_span!("prover.batch_constraints.mle_rounds", phase = "prover").entered();
1570    debug!(%s_deg);
1571    for round in 1..=n_max {
1572        let sp_round_evals = prover.sumcheck_polys_eval(round, r[round - 1]);
1573        let tail_start = prover
1574            .n_per_trace
1575            .iter()
1576            .find_position(|&&n| round as isize > n)
1577            .map(|(i, _)| i)
1578            .unwrap_or(num_traces);
1579        let mut sp_head_zc = vec![SC::EF::ZERO; constraint_degree];
1580        let mut sp_head_logup = vec![SC::EF::ZERO; constraint_degree];
1581        let mut sp_tail = SC::EF::ZERO;
1582        for trace_idx in 0..num_traces {
1583            let zc_idx = 2 * num_traces + trace_idx;
1584            let numer_idx = 2 * trace_idx;
1585            let denom_idx = numer_idx + 1;
1586            if trace_idx < tail_start {
1587                for i in 0..constraint_degree {
1588                    sp_head_zc[i] += mu_pows[zc_idx] * sp_round_evals[zc_idx][i];
1589                    sp_head_logup[i] += mu_pows[numer_idx] * sp_round_evals[numer_idx][i]
1590                        + mu_pows[denom_idx] * sp_round_evals[denom_idx][i];
1591                }
1592            } else {
1593                sp_tail += mu_pows[zc_idx] * sp_round_evals[zc_idx][0]
1594                    + mu_pows[numer_idx] * sp_round_evals[numer_idx][0]
1595                    + mu_pows[denom_idx] * sp_round_evals[denom_idx][0];
1596            }
1597        }
1598        let mut sp_head_evals = vec![SC::EF::ZERO; s_deg];
1599        for i in 0..constraint_degree {
1600            sp_head_evals[i + 1] = prover.eq_ns[round - 1] * sp_head_zc[i]
1601                + prover.eq_sharp_ns[round - 1] * sp_head_logup[i];
1602        }
1603        let xi_cur = prover.xi[l_skip + round - 1];
1604        {
1605            let eq_xi_0 = SC::EF::ONE - xi_cur;
1606            let eq_xi_1 = xi_cur;
1607            sp_head_evals[0] =
1608                (prover.prev_s_eval - eq_xi_1 * sp_head_evals[1] - sp_tail) * eq_xi_0.inverse();
1609        }
1610        let sp_head = UnivariatePoly::lagrange_interpolate(
1611            &(0..s_deg).map(SC::F::from_usize).collect_vec(),
1612            &sp_head_evals,
1613        );
1614        let batch_s = {
1615            let mut coeffs = sp_head.into_coeffs();
1616            coeffs.push(SC::EF::ZERO);
1617            let b = SC::EF::ONE - xi_cur;
1618            let a = xi_cur - b;
1619            for i in (0..s_deg).rev() {
1620                coeffs[i + 1] = a * coeffs[i] + b * coeffs[i + 1];
1621            }
1622            coeffs[0] *= b;
1623            coeffs[1] += sp_tail;
1624            UnivariatePoly::new(coeffs)
1625        };
1626        let batch_s_evals = (1..=s_deg)
1627            .map(|i| batch_s.eval_at_point(SC::EF::from_usize(i)))
1628            .collect_vec();
1629        for &eval in &batch_s_evals {
1630            transcript.observe_ext(eval);
1631        }
1632        sumcheck_round_polys.push(batch_s_evals);
1633
1634        let r_round = transcript.sample_ext();
1635        debug!(%round, %r_round);
1636        r.push(r_round);
1637        prover.prev_s_eval = batch_s.eval_at_point(r_round);
1638
1639        prover.fold_mle_evals(round, r_round);
1640    }
1641    drop(_mle_rounds_span);
1642    assert_eq!(r.len(), n_max + 1);
1643
1644    let column_openings = prover.into_column_openings();
1645
1646    // Observe openings
1647    for (helper, openings) in prover.eval_helpers.iter().zip(column_openings.iter()) {
1648        for (claim, claim_rot) in column_openings_by_rot(&openings[0], helper.needs_next) {
1649            transcript.observe_ext(claim);
1650            transcript.observe_ext(claim_rot);
1651        }
1652    }
1653    for (helper, openings) in prover.eval_helpers.iter().zip(column_openings.iter()) {
1654        for part in openings.iter().skip(1) {
1655            for (claim, claim_rot) in column_openings_by_rot(part, helper.needs_next) {
1656                transcript.observe_ext(claim);
1657                transcript.observe_ext(claim_rot);
1658            }
1659        }
1660    }
1661
1662    let batch_constraint_proof = BatchConstraintProof::<SC> {
1663        numerator_term_per_air,
1664        denominator_term_per_air,
1665        univariate_round_coeffs: s_0_poly.into_coeffs(),
1666        sumcheck_round_polys,
1667        column_openings,
1668    };
1669    let gkr_proof = GkrProof::<SC> {
1670        logup_pow_witness,
1671        q0_claim: frac_sum_proof.fractional_sum.1,
1672        claims_per_layer: frac_sum_proof.claims_per_layer,
1673        sumcheck_polys: frac_sum_proof.sumcheck_polys,
1674    };
1675    Ok((gkr_proof, batch_constraint_proof, r))
1676}