1use 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#[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#[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 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
103fn 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
142fn 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 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 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 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 let values: Vec<EF> = (0..new_height)
188 .into_par_iter()
189 .flat_map(|x| {
190 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 for v in &mut result {
203 *v *= scaling_factor;
204 }
205 result
206 })
207 .collect();
208
209 RowMajorMatrix::new(values, width)
210}
211
212fn 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 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 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 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 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 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 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
307fn 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 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 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 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 let mut node_buf: Vec<<SC::F as Field>::Packing> =
363 Vec::with_capacity(helper.constraints_dag.nodes.len());
364
365 let mut results: Vec<SC::EF> = vec![SC::EF::ZERO; d * skip_sz];
367
368 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 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 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 let evaluator = helper.evaluator_packed(&packed_row_parts);
402 eval_nodes_into(&evaluator, &helper.constraints_dag.nodes, &mut node_buf);
403
404 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 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 let values: Vec<SC::EF> = evals;
443 UnivariatePoly::from_geometric_cosets_evals_idft(RowMajorMatrix::new(values, d), g, g)
444}
445
446struct ViewPair<T> {
451 local: *const T,
452 next: Option<*const T>,
453}
454
455unsafe 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 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
520struct 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#[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
615pub(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 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 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 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 row_parts.push(mat.values[start..start + mat.width].to_vec());
815 }
816 row_parts
817 }
818}
819
820pub(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 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 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 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); vals[i * 3 + 1] = SC::F::from_bool(row_in_period != height - 1); vals[i * 3 + 2] = SC::F::from_bool(row_in_period == height - 1);
1051 }
1053 RowMajorMatrix::new(vals, 3)
1054 })
1055 .collect_vec();
1056
1057 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 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 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 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 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 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 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#[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 let has_interactions = !prover.interactions_layout.sorted_cols.is_empty();
1400 let gkr_input_evals = if !has_interactions {
1401 vec![]
1402 } else {
1403 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 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 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 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 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}