Skip to main content

openvm_stark_backend/prover/
whir.rs

1use std::{iter::once, sync::Arc};
2
3use itertools::Itertools;
4use p3_dft::{Radix2DitParallel, TwoAdicSubgroupDft};
5use p3_field::{ExtensionField, Field, PrimeCharacteristicRing, TwoAdicField};
6use p3_maybe_rayon::prelude::*;
7use p3_util::log2_strict_usize;
8use tracing::instrument;
9
10use crate::{
11    poly_common::Squarable,
12    proof::{MerkleProof, WhirProof},
13    prover::{
14        error::WhirProverError,
15        poly::{eval_to_coeff_rs_message, evals_eq_hypercube, evals_mobius_eq_hypercube, Mle},
16        stacked_pcs::{MerkleTree, StackedPcsData},
17        ColMajorMatrix, CpuColMajorBackend, MatrixDimensions, ProverBackend, ReferenceDevice,
18    },
19    FiatShamirTranscript, StarkProtocolConfig, WhirConfig,
20};
21
22pub trait WhirProver<SC: StarkProtocolConfig, PB: ProverBackend, PD, TS> {
23    type Error;
24
25    /// Prove the WHIR protocol for a collection of MLE polynomials \hat{q}_j, each in n variables,
26    /// at a single vector `u \in \Fext^n`.
27    ///
28    /// This means applying WHIR with weight polynomial
29    /// `\hat{w}(Z, \vec X) = Z * mobius_eq_poly(u)(\vec X)`, where `mobius_eq_poly(u)` is the
30    /// Möbius-adjusted equality polynomial for eval-to-coeff RS encoding.
31    ///
32    /// The matrices in `common_main_pcs_data` and `pre_cached_pcs_data_per_commit` must all have
33    /// the same height.
34    fn prove_whir(
35        &self,
36        transcript: &mut TS,
37        common_main_pcs_data: PB::PcsData,
38        pre_cached_pcs_data_per_commit: Vec<Arc<PB::PcsData>>,
39        u_cube: &[PB::Challenge],
40    ) -> Result<WhirProof<SC>, Self::Error>;
41}
42
43impl<SC, TS> WhirProver<SC, CpuColMajorBackend<SC>, ReferenceDevice<SC>, TS> for ReferenceDevice<SC>
44where
45    SC: StarkProtocolConfig,
46    SC::F: TwoAdicField + Ord,
47    SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
48    TS: FiatShamirTranscript<SC>,
49{
50    type Error = WhirProverError;
51
52    #[instrument(level = "info", skip_all)]
53    fn prove_whir(
54        &self,
55        transcript: &mut TS,
56        common_main_pcs_data: StackedPcsData<SC::F, SC::Digest>,
57        pre_cached_pcs_data_per_commit: Vec<Arc<StackedPcsData<SC::F, SC::Digest>>>,
58        u_cube: &[SC::EF],
59    ) -> Result<WhirProof<SC>, WhirProverError> {
60        let params = self.params();
61        let committed_mats = once(&common_main_pcs_data)
62            .chain(pre_cached_pcs_data_per_commit.iter().map(|d| d.as_ref()))
63            .map(|d| (&d.matrix, &d.tree))
64            .collect_vec();
65        prove_whir_opening::<SC, _>(
66            transcript,
67            self.config().hasher(),
68            params.l_skip,
69            params.log_blowup,
70            &params.whir,
71            &committed_mats,
72            u_cube,
73        )
74    }
75}
76
77#[allow(clippy::too_many_arguments, clippy::type_complexity)]
78pub fn prove_whir_opening<SC, TS>(
79    transcript: &mut TS,
80    hasher: &SC::Hasher,
81    l_skip: usize,
82    log_blowup: usize,
83    whir_params: &WhirConfig,
84    committed_mats: &[(&ColMajorMatrix<SC::F>, &MerkleTree<SC::F, SC::Digest>)],
85    u: &[SC::EF],
86) -> Result<WhirProof<SC>, WhirProverError>
87where
88    SC: StarkProtocolConfig,
89    SC::F: TwoAdicField + Ord,
90    SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
91    TS: FiatShamirTranscript<SC>,
92{
93    // Proof-of-work grinding before μ batching challenge.
94    // This amplifies soundness of the initial batching step.
95    let mu_pow_witness = transcript.grind(whir_params.mu_pow_bits);
96
97    // Sample randomness for algebraic batching.
98    // We batch the codewords for \hat{q}_j together _before_ applying WHIR.
99    let mu = transcript.sample_ext();
100    let total_width = committed_mats.iter().map(|(mat, _)| mat.width()).sum();
101    let mu_powers = mu.powers().take(total_width).collect_vec();
102
103    let height = committed_mats[0].0.height();
104    debug_assert!(committed_mats.iter().all(|(mat, _)| mat.height() == height));
105    let mut m = log2_strict_usize(height);
106
107    let k_whir = whir_params.k;
108    let num_whir_rounds = whir_params.num_whir_rounds();
109    let num_sumcheck_rounds = whir_params.num_sumcheck_rounds();
110
111    let mles: Vec<Vec<SC::F>> = committed_mats
112        .par_iter()
113        .flat_map(|(mat, _)| {
114            mat.par_columns().map(|col| {
115                // Convert column evaluations directly into eval-to-coeff RS coefficients, then
116                // interpret them as MLE coefficients of HatF and compute HatF hypercube evals.
117                let mut x = eval_to_coeff_rs_message(l_skip, col);
118                Mle::coeffs_to_evals_inplace(&mut x);
119                x
120            })
121        })
122        .collect();
123
124    // The evaluations of `\hat{f}` in the current WHIR round on the hypercube `H_m`.
125    let mut f_evals: Vec<_> = (0..1 << m)
126        .into_par_iter()
127        .map(|i| {
128            mles.iter()
129                .zip(mu_powers.iter())
130                .fold(SC::EF::ZERO, |acc, (mle_j, mu_j)| acc + *mu_j * mle_j[i])
131        })
132        .collect();
133
134    // We assume `\hat{w}` in a WHIR round is always multilinear and maintain its
135    // evaluations on `H_m`.
136    let mut w_evals = evals_mobius_eq_hypercube(u);
137
138    let mut whir_sumcheck_polys: Vec<[SC::EF; 2]> = Vec::with_capacity(num_sumcheck_rounds);
139    let mut codeword_commits = vec![];
140    let mut ood_values = vec![];
141    // per commitment, per whir query, per column
142    let mut initial_round_opened_rows: Vec<Vec<Vec<Vec<SC::F>>>> =
143        vec![vec![]; committed_mats.len()];
144    let mut initial_round_merkle_proofs: Vec<Vec<MerkleProof<SC::Digest>>> =
145        vec![vec![]; committed_mats.len()];
146    let mut codeword_opened_values: Vec<Vec<Vec<SC::EF>>> = Vec::with_capacity(num_whir_rounds - 1);
147    let mut codeword_merkle_proofs: Vec<Vec<MerkleProof<SC::Digest>>> =
148        Vec::with_capacity(num_whir_rounds - 1);
149    let mut folding_pow_witnesses = Vec::with_capacity(num_sumcheck_rounds);
150    let mut query_phase_pow_witnesses = Vec::with_capacity(num_whir_rounds);
151    let mut rs_tree = None;
152    let mut log_rs_domain_size = m + log_blowup;
153    let mut final_poly = None;
154    for (whir_round, round_params) in whir_params.rounds.iter().enumerate() {
155        let is_last_round = whir_round == num_whir_rounds - 1;
156        // Run k_whir rounds of sumcheck on `sum_{x in H_m} \hat{w}(\hat{f}(x), x)`
157        for round in 0..k_whir {
158            debug_assert_eq!(f_evals.len(), 1 << (m - round));
159
160            // \hat{f} * eq has degree 2
161            let s_deg = 2;
162            let s_evals = (1..=s_deg)
163                .map(|x| {
164                    let x = SC::F::from_usize(x);
165                    let hypercube_dim = m - round - 1;
166                    (0..(1usize << hypercube_dim))
167                        .map(|y| {
168                            let f_0 = f_evals[y << 1];
169                            let f_1 = f_evals[(y << 1) + 1];
170                            let f_x = f_0 + (f_1 - f_0) * x;
171                            let w_0 = w_evals[y << 1];
172                            let w_1 = w_evals[(y << 1) + 1];
173                            let w_x = w_0 + (w_1 - w_0) * x;
174                            f_x * w_x
175                        })
176                        .fold(SC::EF::ZERO, |acc, x| acc + x)
177                })
178                .collect_vec();
179
180            for &eval in &s_evals {
181                transcript.observe_ext(eval);
182            }
183            whir_sumcheck_polys.push(
184                s_evals
185                    .try_into()
186                    .map_err(|_| WhirProverError::TryIntoFailed)?,
187            );
188
189            folding_pow_witnesses.push(transcript.grind(whir_params.folding_pow_bits));
190            // Folding randomness
191            let alpha = transcript.sample_ext();
192
193            // Fold the evaluations
194            let half = f_evals.len() / 2;
195            for y in 0..half {
196                let eval_0 = f_evals[y << 1];
197                let eval_1 = f_evals[(y << 1) + 1];
198                // Linear interpolation at r_round
199                f_evals[y] = eval_0 + alpha * (eval_1 - eval_0);
200
201                let eval_0 = w_evals[y << 1];
202                let eval_1 = w_evals[(y << 1) + 1];
203                w_evals[y] = eval_0 + alpha * (eval_1 - eval_0);
204            }
205            f_evals.truncate(half);
206            w_evals.truncate(half);
207        }
208        // Define g^ = f^(alpha, \cdot) and send matrix commit of RS(g^)
209        // f_evals is the evaluations of f^(alpha, \cdot) on hypercube
210        let g_mle = Mle::from_evaluations(&f_evals);
211        let (g_tree, z_0) = if !is_last_round {
212            let dft = Radix2DitParallel::default();
213            let mut g_coeffs = g_mle.coeffs().to_vec();
214            debug_assert_eq!(g_coeffs.len(), 1 << (m - k_whir));
215            g_coeffs.resize(1 << (log_rs_domain_size - 1), SC::EF::ZERO);
216            // `g: \mathcal{L}^{(2)} \to \mathbb F`
217            let g_rs = dft.dft(g_coeffs);
218            let g_tree = MerkleTree::new(hasher, ColMajorMatrix::new(g_rs, 1), 1 << k_whir)?;
219            let g_commit = g_tree.root()?;
220            transcript.observe_commit(g_commit);
221            codeword_commits.push(g_commit);
222
223            let z_0 = transcript.sample_ext();
224            let z_0_vec = z_0.exp_powers_of_2().take(m - k_whir).collect_vec();
225            let g_opened_value = g_mle.eval_at_point(&z_0_vec);
226            transcript.observe_ext(g_opened_value);
227            ood_values.push(g_opened_value);
228
229            (Some(g_tree), Some(z_0))
230        } else {
231            let coeffs = g_mle.into_coeffs();
232            for coeff in &coeffs {
233                transcript.observe_ext(*coeff);
234            }
235            final_poly = Some(coeffs);
236            (None, None)
237        };
238
239        // omega is generator of RS domain `\mathcal{L}^{(2^k)}`
240        let omega = SC::F::two_adic_generator(log_rs_domain_size - k_whir);
241        let num_queries = round_params.num_queries;
242        let mut query_indices = Vec::with_capacity(num_queries);
243        query_phase_pow_witnesses.push(transcript.grind(whir_params.query_phase_pow_bits));
244        // Sample query indices first
245        for _ in 0..num_queries {
246            // This is the index of the leaf in the Merkle tree
247            let index = transcript.sample_bits(log_rs_domain_size - k_whir);
248            query_indices.push(index as usize);
249        }
250        let mut zs = Vec::with_capacity(num_queries);
251        if !is_last_round {
252            codeword_opened_values.push(vec![]);
253            codeword_merkle_proofs.push(vec![]);
254        }
255        for (query_idx, index) in query_indices.into_iter().enumerate() {
256            let z_i = omega.exp_u64(index as u64);
257            // Get merkle proofs for in-domain samples necessary to evaluate Fold(f, \vec
258            // \alpha)(z_i)
259            zs.push(z_i);
260
261            let depth = log_rs_domain_size.saturating_sub(k_whir);
262            // Row openings are different between first WHIR round (width > 1) and other rounds
263            // (width = 1):
264            // NOTE: merkle proof is deterministic from the index and merkle root, so the opened_row
265            // and merkle proof are both hinted and not observed by the transcript.
266            if whir_round == 0 {
267                #[allow(clippy::needless_range_loop)]
268                for com_idx in 0..committed_mats.len() {
269                    debug_assert_eq!(initial_round_merkle_proofs[com_idx].len(), query_idx);
270                    let tree = &committed_mats[com_idx].1;
271                    let tree_height = tree.backing_matrix.height();
272                    let expected = 1 << log_rs_domain_size;
273                    if tree_height != expected {
274                        return Err(WhirProverError::TreeHeightMismatch {
275                            tree_height,
276                            expected,
277                        });
278                    }
279                    let opened_rows = tree.get_opened_rows(index)?;
280                    initial_round_opened_rows[com_idx].push(opened_rows);
281                    debug_assert_eq!(tree.proof_depth(), depth);
282                    let proof = tree.query_merkle_proof(index)?;
283                    debug_assert_eq!(proof.len(), depth);
284                    initial_round_merkle_proofs[com_idx].push(proof);
285                }
286            } else {
287                let tree: &MerkleTree<SC::EF, SC::Digest> =
288                    rs_tree.as_ref().ok_or(WhirProverError::RsTreeNone)?;
289                let width = tree.backing_matrix.width();
290                if width != 1 {
291                    return Err(WhirProverError::TreeWidthNotOne { width });
292                }
293                let opened_rows = tree
294                    .get_opened_rows(index)?
295                    .into_iter()
296                    .flatten()
297                    .collect_vec();
298                codeword_opened_values[whir_round - 1].push(opened_rows);
299                debug_assert_eq!(tree.proof_depth(), depth);
300                let proof = tree.query_merkle_proof(index)?;
301                debug_assert_eq!(proof.len(), depth);
302                codeword_merkle_proofs[whir_round - 1].push(proof);
303            }
304        }
305        rs_tree = g_tree;
306
307        // We still sample on the last round to match the verifier, who uses a
308        // final gamma to unify some logic. But we do not need to update
309        // `w_evals`.
310        let gamma = transcript.sample_ext();
311
312        if !is_last_round {
313            // Update \hat{w}
314            w_evals_accumulate::<SC::EF, SC::EF>(
315                &mut w_evals,
316                z_0.ok_or(WhirProverError::Z0None)?,
317                gamma,
318            );
319            for (z_i, gamma_pow) in zs.into_iter().zip(gamma.powers().skip(2)) {
320                w_evals_accumulate::<SC::F, SC::EF>(&mut w_evals, z_i, gamma_pow);
321            }
322        }
323
324        m -= k_whir;
325        log_rs_domain_size -= 1;
326    }
327
328    Ok(WhirProof::<SC> {
329        mu_pow_witness,
330        whir_sumcheck_polys,
331        codeword_commits,
332        ood_values,
333        folding_pow_witnesses,
334        query_phase_pow_witnesses,
335        initial_round_opened_rows,
336        initial_round_merkle_proofs,
337        codeword_opened_values,
338        codeword_merkle_proofs,
339        final_poly: final_poly.ok_or(WhirProverError::FinalPolyNone)?,
340    })
341}
342
343/// Given hypercube evaluations `w_evals` of `\hat{w}` on `H_t`, this updates the evaluations
344/// in place to be the evaluations of `\hat{w}'(x) = \hat{w}(x) + γ * eq(x, pow(z))`.
345fn w_evals_accumulate<F: Field, EF: ExtensionField<F>>(w_evals: &mut [EF], z: F, gamma: EF) {
346    let dim = log2_strict_usize(w_evals.len());
347    let z_pows = z.exp_powers_of_2().take(dim).collect_vec();
348    let evals = evals_eq_hypercube(&z_pows);
349    for (w, x) in w_evals.iter_mut().zip(evals.into_iter()) {
350        *w += gamma * x;
351    }
352}