Skip to main content

openvm_stark_backend/prover/logup_zerocheck/
single.rs

1//! Single AIR constraint evaluation helpers
2
3use 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
17/// For a single AIR
18pub struct EvalHelper<'a, F> {
19    /// AIR constraints
20    pub constraints_dag: &'a SymbolicExpressionDag<F>,
21    /// Interactions
22    pub interactions: Vec<SymbolicInteraction<F>>,
23    pub public_values: Vec<F>,
24    pub preprocessed_trace: Option<StridedColMajorMatrixView<'a, F>>,
25    // PERF: skip rotation if vk dictates it is never used
26    pub needs_next: bool,
27    pub constraint_degree: u8,
28}
29
30impl<'a, F: TwoAdicField> EvalHelper<'a, F> {
31    /// Returns list of (ref to column-major matrix, is_rot) pairs in the order:
32    /// - (if has_preprocessed) (preprocessed, false), (preprocessed, true)
33    /// - (cached_0, false), (cached_0, true), ..., (cached_{m-1}, false), (cached_{m-1}, true)
34    /// - (common, false), (common, true)
35    ///
36    /// Note: currently every matrix returns both non-rotated and rotated versions. This will change
37    /// in the future for perf.
38    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    /// See [Self::evaluator].
76    // Assumes that `z[0] != 1` or `omega_D^{-1}` to avoid handling division by zero.
77    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    /// See [Self::evaluator].
91    ///
92    /// Returns sum of ordered list of `interactions`, weighted by `eq(\xi_3, b_{T,\hat\sigma})`
93    /// terms as (numerator, denominator) pair.
94    ///
95    /// Note: the denominator does not include the `alpha` term.
96    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        // PERF[jpw]: no need to collect the vec, but I ran into a lifetime issue returning iterator
107        // in `eval_interactions`
108        let interaction_evals = self.eval_interactions(row_parts, beta_pows);
109        let mut numer = EF::ZERO;
110        let mut denom = EF::ZERO; // without alpha term
111        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    // `row_parts` should have separate Vec in following order:
148    // - selectors [is_first_row, is_transition, is_last_row]
149    // - (if has_preprocessed) preprocessed
150    // - (if has_preprocessed) preprocessed_rot
151    // - cached_0
152    // - cached_0_rot
153    // - ...
154    // - common
155    // - common_rot
156    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}