openvm_stark_backend/prover/logup_zerocheck/
single.rs1use std::iter::zip;
4
5use itertools::Itertools;
6use p3_field::{ExtensionField, TwoAdicField};
7
8use crate::{
9 air_builders::symbolic::{symbolic_expression::SymbolicEvaluator, SymbolicExpressionDag},
10 interaction::SymbolicInteraction,
11 prover::{
12 logup_zerocheck::evaluator::{ProverConstraintEvaluator, ViewPair},
13 AirProvingContext, ColMajorMatrix, ProverBackend, StridedColMajorMatrixView,
14 },
15};
16
17pub struct EvalHelper<'a, F> {
19 pub constraints_dag: &'a SymbolicExpressionDag<F>,
21 pub interactions: Vec<SymbolicInteraction<F>>,
23 pub public_values: Vec<F>,
24 pub preprocessed_trace: Option<StridedColMajorMatrixView<'a, F>>,
25 pub needs_next: bool,
27 pub constraint_degree: u8,
28}
29
30impl<'a, F: TwoAdicField> EvalHelper<'a, F> {
31 pub fn view_mats<PB>(
39 &self,
40 ctx: &'a AirProvingContext<PB>,
41 ) -> Vec<(StridedColMajorMatrixView<'a, F>, bool)>
42 where
43 PB: ProverBackend<Val = F, Matrix = ColMajorMatrix<F>>,
44 {
45 let base_mats = usize::from(self.has_preprocessed()) + 1 + ctx.cached_mains.len();
46 let mut mats = Vec::with_capacity(if self.needs_next {
47 2 * base_mats
48 } else {
49 base_mats
50 });
51 if let Some(mat) = self.preprocessed_trace {
52 mats.push((mat, false));
53 if self.needs_next {
54 mats.push((mat, true));
55 }
56 }
57 for cd in ctx.cached_mains.iter() {
58 let trace_view: StridedColMajorMatrixView<'a, F> = cd.trace.as_view().into();
59 mats.push((trace_view, false));
60 if self.needs_next {
61 mats.push((trace_view, true));
62 }
63 }
64 mats.push((ctx.common_main.as_view().into(), false));
65 if self.needs_next {
66 mats.push((ctx.common_main.as_view().into(), true));
67 }
68 mats
69 }
70
71 pub fn has_preprocessed(&self) -> bool {
72 self.preprocessed_trace.is_some()
73 }
74
75 pub fn acc_constraints<FF: ExtensionField<F>, EF: ExtensionField<FF>>(
78 &self,
79 row_parts: &[Vec<FF>],
80 lambda_pows: &[EF],
81 ) -> EF {
82 let evaluator = self.evaluator(row_parts);
83 let nodes = evaluator.eval_nodes(&self.constraints_dag.nodes);
84 zip(lambda_pows, &self.constraints_dag.constraint_idx)
85 .fold(EF::ZERO, |acc, (&lambda_pow, &idx)| {
86 acc + lambda_pow * nodes[idx]
87 })
88 }
89
90 pub fn acc_interactions<FF, EF>(
97 &self,
98 row_parts: &[Vec<FF>],
99 beta_pows: &[EF],
100 eq_3bs: &[EF],
101 ) -> [EF; 2]
102 where
103 FF: ExtensionField<F>,
104 EF: ExtensionField<FF> + ExtensionField<F>,
105 {
106 let interaction_evals = self.eval_interactions(row_parts, beta_pows);
109 let mut numer = EF::ZERO;
110 let mut denom = EF::ZERO; for (&eq_3b, eval) in zip(eq_3bs, interaction_evals) {
112 numer += eq_3b * eval.0;
113 denom += eq_3b * eval.1;
114 }
115 [numer, denom]
116 }
117
118 pub fn eval_interactions<FF, EF>(
119 &self,
120 row_parts: &[Vec<FF>],
121 beta_pows: &[EF],
122 ) -> Vec<(FF, EF)>
123 where
124 FF: ExtensionField<F>,
125 EF: ExtensionField<FF> + ExtensionField<F>,
126 {
127 let evaluator = self.evaluator(row_parts);
128 self.interactions
129 .iter()
130 .map(|interaction| {
131 let b = F::from_u32(interaction.bus_index as u32 + 1);
132 let msg_len = interaction.message.len();
133 debug_assert!(msg_len <= beta_pows.len());
134 let denom = zip(&interaction.message, beta_pows).fold(
135 beta_pows[msg_len] * b,
136 |h_beta, (msg_j, &beta_j)| {
137 let msg_j_eval = evaluator.eval_expr(msg_j);
138 h_beta + beta_j * msg_j_eval
139 },
140 );
141 let numer = evaluator.eval_expr(&interaction.count);
142 (numer, denom)
143 })
144 .collect()
145 }
146
147 fn evaluator<FF: ExtensionField<F>>(
157 &self,
158 row_parts: &[Vec<FF>],
159 ) -> ProverConstraintEvaluator<'_, F, FF> {
160 let sels = &row_parts[0];
161 let mut view_pairs = if self.needs_next {
162 let mut chunks = row_parts[1..].chunks_exact(2);
163 let pairs = chunks
164 .by_ref()
165 .map(|pair| ViewPair::new(&pair[0], Some(&pair[1][..])))
166 .collect_vec();
167 debug_assert!(chunks.remainder().is_empty());
168 pairs
169 } else {
170 row_parts[1..]
171 .iter()
172 .map(|part| ViewPair::new(part, None))
173 .collect_vec()
174 };
175 let mut preprocessed = None;
176 if self.has_preprocessed() {
177 preprocessed = Some(view_pairs.remove(0));
178 }
179 ProverConstraintEvaluator {
180 preprocessed,
181 partitioned_main: view_pairs,
182 is_first_row: sels[0],
183 is_transition: sels[1],
184 is_last_row: sels[2],
185 public_values: &self.public_values,
186 }
187 }
188}