openvm_stark_backend/prover/
whir.rs1use 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 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 ¶ms.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 let mu_pow_witness = transcript.grind(whir_params.mu_pow_bits);
96
97 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 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 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 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 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 for round in 0..k_whir {
158 debug_assert_eq!(f_evals.len(), 1 << (m - round));
159
160 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 let alpha = transcript.sample_ext();
192
193 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 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 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 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 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 for _ in 0..num_queries {
246 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 zs.push(z_i);
260
261 let depth = log_rs_domain_size.saturating_sub(k_whir);
262 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 let gamma = transcript.sample_ext();
311
312 if !is_last_round {
313 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
343fn 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}