1use core::{cmp, iter::zip, ops::Range};
2use std::sync::Arc;
3
4use itertools::{izip, Itertools};
5use openvm_circuit_primitives::encoder::Encoder;
6use openvm_cpu_backend::CpuBackend;
7use openvm_stark_backend::{
8 keygen::types::MultiStarkVerifyingKey,
9 p3_maybe_rayon::prelude::*,
10 poly_common::{eval_mle_evals_at_point, interpolate_quadratic_at_012, Squarable},
11 proof::{Proof, WhirProof},
12 prover::AirProvingContext,
13 AirRef, FiatShamirTranscript, StarkProtocolConfig, SystemParams, TranscriptHistory,
14};
15use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, CHUNK, EF, F};
16use p3_field::{Field, PrimeCharacteristicRing, PrimeField32, TwoAdicField};
17use p3_matrix::dense::RowMajorMatrix;
18use strum::{EnumCount, EnumDiscriminants};
19
20#[cfg(feature = "cuda")]
21use crate::primitives::exp_bits_len::ExpBitsLenTraceGenerator as GpuExpBitsLenTraceGenerator;
22use crate::{
23 primitives::exp_bits_len::ExpBitsLenCpuTraceGenerator,
24 system::{
25 AirModule, BusIndexManager, BusInventory, GlobalCtxCpu, Preflight, TraceGenModule,
26 WhirPreflight,
27 },
28 tracegen::{ModuleChip, RowMajorChip, StandardTracegenCtx},
29 utils::{pow_observe_sample, FlattenedLayout, FlattenedVec},
30 whir::{
31 bus::{
32 FinalPolyFoldingBus, FinalPolyMleEvalBus, FinalPolyQueryEvalBus, VerifyQueriesBus,
33 VerifyQueryBus, WhirAlphaBus, WhirEqAlphaUBus, WhirFinalPolyBus, WhirFoldingBus,
34 WhirGammaBus, WhirQueryBus, WhirSumcheckBus,
35 },
36 final_poly_mle_eval::FinalPolyMleEvalAir,
37 final_poly_query_eval::FinalPolyQueryEvalAir,
38 folding::{FoldRecord, WhirFoldingAir},
39 initial_opened_values::InitialOpenedValuesAir,
40 non_initial_opened_values::NonInitialOpenedValuesAir,
41 query::WhirQueryAir,
42 sumcheck::SumcheckAir,
43 whir_round::WhirRoundAir,
44 },
45};
46
47mod bus;
48mod final_poly_mle_eval;
49mod final_poly_query_eval;
50pub mod folding;
51mod initial_opened_values;
52mod non_initial_opened_values;
53mod query;
54mod sumcheck;
55mod whir_round;
56
57pub(crate) fn num_queries_per_round(params: &SystemParams) -> Vec<usize> {
58 params
59 .whir
60 .rounds
61 .iter()
62 .map(|round| round.num_queries)
63 .collect()
64}
65
66pub(crate) fn whir_round_encoder(num_rounds: usize) -> Encoder {
67 Encoder::new(num_rounds.max(2), 2, false)
69}
70
71#[inline]
72fn eval_final_poly_at_u(final_poly: &[EF], u_tail: &[EF]) -> EF {
73 let mut evals = final_poly.to_vec();
74 eval_mle_evals_at_point(&mut evals, u_tail)
75}
76
77pub(in crate::whir) type PerProofIdx = (usize, usize);
78pub(in crate::whir) type QueryIdx = (usize, usize, usize);
79
80#[derive(Clone, Debug)]
81pub(in crate::whir) struct PerProofLayout {
82 num_proofs: usize,
83 items_per_proof: usize,
84}
85
86impl PerProofLayout {
87 pub(in crate::whir) fn new(num_proofs: usize, items_per_proof: usize) -> Self {
88 Self {
89 num_proofs,
90 items_per_proof,
91 }
92 }
93
94 #[inline]
95 pub(in crate::whir) fn items_per_proof(&self) -> usize {
96 self.items_per_proof
97 }
98}
99
100impl FlattenedLayout for PerProofLayout {
101 type Index = PerProofIdx;
102
103 #[inline]
104 fn len(&self) -> usize {
105 self.num_proofs * self.items_per_proof
106 }
107
108 #[inline]
109 fn offset(&self, idx: Self::Index) -> usize {
110 let (proof_idx, item_idx) = idx;
111 debug_assert!(
112 proof_idx < self.num_proofs,
113 "proof index out of bounds: {proof_idx} >= {}",
114 self.num_proofs
115 );
116 debug_assert!(
117 item_idx < self.items_per_proof,
118 "index out of bounds: {item_idx} >= {}",
119 self.items_per_proof
120 );
121 proof_idx * self.items_per_proof + item_idx
122 }
123}
124
125#[derive(Clone, Debug)]
126pub(in crate::whir) struct VariablePerProofLayout {
127 proof_offsets: Vec<usize>,
128}
129
130impl VariablePerProofLayout {
131 pub(in crate::whir) fn new(per_proof_items: impl IntoIterator<Item = usize>) -> Self {
132 let mut proof_offsets = vec![0];
133 let mut total = 0usize;
134 for items in per_proof_items {
135 total += items;
136 proof_offsets.push(total);
137 }
138 Self { proof_offsets }
139 }
140}
141
142impl FlattenedLayout for VariablePerProofLayout {
143 type Index = PerProofIdx;
144
145 #[inline]
146 fn len(&self) -> usize {
147 *self.proof_offsets.last().unwrap_or(&0)
148 }
149
150 #[inline]
151 fn offset(&self, idx: Self::Index) -> usize {
152 let (proof_idx, item_idx) = idx;
153 debug_assert!(
154 proof_idx + 1 < self.proof_offsets.len(),
155 "proof index out of bounds: {proof_idx} >= {}",
156 self.proof_offsets.len().saturating_sub(1)
157 );
158 let proof_start = self.proof_offsets[proof_idx];
159 let proof_end = self.proof_offsets[proof_idx + 1];
160 debug_assert!(
161 item_idx < proof_end - proof_start,
162 "index out of bounds: {} >= {} for proof {}",
163 item_idx,
164 proof_end - proof_start,
165 proof_idx
166 );
167 proof_start + item_idx
168 }
169}
170
171#[derive(Clone, Debug)]
172pub(in crate::whir) struct WhirQueryLayout {
173 num_proofs: usize,
174 query_offsets: Vec<usize>,
175}
176
177impl WhirQueryLayout {
178 pub(in crate::whir) fn new(num_proofs: usize, num_queries_per_round: &[usize]) -> Self {
179 let mut query_offsets = Vec::with_capacity(num_queries_per_round.len() + 1);
180 query_offsets.push(0);
181 for &num_queries in num_queries_per_round {
182 query_offsets.push(query_offsets.last().copied().unwrap() + num_queries);
183 }
184 Self {
185 num_proofs,
186 query_offsets,
187 }
188 }
189
190 #[inline]
191 pub(in crate::whir) fn num_rounds(&self) -> usize {
192 debug_assert!(!self.query_offsets.is_empty());
193 self.query_offsets.len() - 1
194 }
195
196 #[inline]
197 pub(in crate::whir) fn round_query_range(&self, whir_round: usize) -> Range<usize> {
198 debug_assert!(
199 whir_round + 1 < self.query_offsets.len(),
200 "WHIR round out of bounds: {whir_round} >= {}",
201 self.query_offsets.len().saturating_sub(1)
202 );
203 self.query_offsets[whir_round]..self.query_offsets[whir_round + 1]
204 }
205
206 #[inline]
207 pub(in crate::whir) fn round_num_queries(&self, whir_round: usize) -> usize {
208 self.round_query_range(whir_round).len()
209 }
210
211 #[inline]
212 pub(in crate::whir) fn iter_round_query_ranges(
213 &self,
214 ) -> impl Iterator<Item = (usize, Range<usize>)> + '_ {
215 (0..self.num_rounds())
216 .map(move |whir_round| (whir_round, self.round_query_range(whir_round)))
217 }
218
219 #[inline]
220 pub(in crate::whir) fn round_and_query_idx(&self, proof_query_idx: usize) -> (usize, usize) {
221 debug_assert!(
222 proof_query_idx < self.queries_per_proof(),
223 "proof query index out of bounds: {proof_query_idx} >= {}",
224 self.queries_per_proof()
225 );
226 let whir_round =
227 self.query_offsets[1..].partition_point(|&offset| offset <= proof_query_idx);
228 let query_idx = proof_query_idx - self.query_offsets[whir_round];
229 (whir_round, query_idx)
230 }
231
232 #[inline]
233 pub(in crate::whir) fn queries_per_proof(&self) -> usize {
234 *self.query_offsets.last().unwrap_or(&0)
235 }
236
237 #[cfg(feature = "cuda")]
239 #[inline]
240 pub(in crate::whir) fn query_offsets(&self) -> &[usize] {
241 &self.query_offsets
242 }
243}
244
245impl FlattenedLayout for WhirQueryLayout {
246 type Index = QueryIdx;
247
248 #[inline]
249 fn len(&self) -> usize {
250 self.num_proofs * self.queries_per_proof()
251 }
252
253 #[inline]
254 fn offset(&self, idx: Self::Index) -> usize {
255 let (proof_idx, whir_round, query_idx) = idx;
256 debug_assert!(
257 proof_idx < self.num_proofs,
258 "proof index out of bounds: {proof_idx} >= {}",
259 self.num_proofs
260 );
261 debug_assert!(
262 whir_round + 1 < self.query_offsets.len(),
263 "WHIR round out of bounds: {whir_round} >= {}",
264 self.query_offsets.len().saturating_sub(1)
265 );
266 let proof_start = proof_idx * self.queries_per_proof();
267 let round_start = proof_start + self.query_offsets[whir_round];
268 let round_end = proof_start + self.query_offsets[whir_round + 1];
269 debug_assert!(
270 query_idx < round_end - round_start,
271 "query index out of bounds: {} >= {} for proof {}, round {}",
272 query_idx,
273 round_end - round_start,
274 proof_idx,
275 whir_round
276 );
277 round_start + query_idx
278 }
279}
280
281pub(in crate::whir) type CodewordAccsIdx = (usize, usize, usize, usize, usize);
282
283#[derive(Clone, Debug)]
286pub(in crate::whir) struct CodewordAccsLayout {
287 num_queries: usize,
288 num_cosets: usize,
289 rows_per_proof_offsets: Vec<usize>,
290 commits_per_proof_offsets: Vec<usize>,
291 stacking_chunks_offsets: Vec<usize>,
292 stacking_widths_offsets: Vec<usize>,
293}
294
295impl CodewordAccsLayout {
296 pub(in crate::whir) fn new(
297 num_queries: usize,
298 num_cosets: usize,
299 per_proof_commit_widths: &[Vec<usize>],
300 ) -> Self {
301 let num_proofs = per_proof_commit_widths.len();
302 let mut rows_per_proof_offsets = Vec::with_capacity(num_proofs + 1);
303 let mut commits_per_proof_offsets = Vec::with_capacity(num_proofs + 1);
304 let total_commits: usize = per_proof_commit_widths.iter().map(Vec::len).sum();
305 let mut stacking_chunks_offsets = Vec::with_capacity(total_commits + 1);
306 let mut stacking_widths_offsets = Vec::with_capacity(total_commits + 1);
307
308 rows_per_proof_offsets.push(0);
309 commits_per_proof_offsets.push(0);
310 stacking_chunks_offsets.push(0);
311 stacking_widths_offsets.push(0);
312
313 for per_commit_widths in per_proof_commit_widths {
314 let mut total_chunks_for_proof = 0usize;
315 for &width in per_commit_widths {
316 let chunks = width.div_ceil(CHUNK);
317 total_chunks_for_proof += chunks;
318 stacking_chunks_offsets.push(*stacking_chunks_offsets.last().unwrap() + chunks);
319 stacking_widths_offsets.push(*stacking_widths_offsets.last().unwrap() + width);
320 }
321 rows_per_proof_offsets.push(
322 *rows_per_proof_offsets.last().unwrap()
323 + num_queries * num_cosets * total_chunks_for_proof,
324 );
325 commits_per_proof_offsets
326 .push(*commits_per_proof_offsets.last().unwrap() + per_commit_widths.len());
327 }
328 Self {
329 num_queries,
330 num_cosets,
331 rows_per_proof_offsets,
332 commits_per_proof_offsets,
333 stacking_chunks_offsets,
334 stacking_widths_offsets,
335 }
336 }
337
338 #[inline]
339 pub(in crate::whir) fn num_proofs(&self) -> usize {
340 self.rows_per_proof_offsets.len().saturating_sub(1)
341 }
342
343 #[inline]
344 pub(in crate::whir) fn total_chunks(&self, proof_idx: usize) -> usize {
345 let commit_start = self.commits_per_proof_offsets[proof_idx];
346 let commit_end = self.commits_per_proof_offsets[proof_idx + 1];
347 self.stacking_chunks_offsets[commit_end] - self.stacking_chunks_offsets[commit_start]
348 }
349
350 #[cfg(feature = "cuda")]
351 #[inline]
352 pub(in crate::whir) fn rows_per_proof_offsets(&self) -> &[usize] {
353 &self.rows_per_proof_offsets
354 }
355
356 #[cfg(feature = "cuda")]
357 #[inline]
358 pub(in crate::whir) fn commits_per_proof_offsets(&self) -> &[usize] {
359 &self.commits_per_proof_offsets
360 }
361
362 #[cfg(feature = "cuda")]
363 #[inline]
364 pub(in crate::whir) fn stacking_chunks_offsets(&self) -> &[usize] {
365 &self.stacking_chunks_offsets
366 }
367
368 #[cfg(feature = "cuda")]
369 #[inline]
370 pub(in crate::whir) fn stacking_widths_offsets(&self) -> &[usize] {
371 &self.stacking_widths_offsets
372 }
373
374 #[inline]
375 pub(in crate::whir) fn total_width(&self, proof_idx: usize) -> usize {
376 let commit_start = self.commits_per_proof_offsets[proof_idx];
377 let commit_end = self.commits_per_proof_offsets[proof_idx + 1];
378 self.stacking_widths_offsets[commit_end] - self.stacking_widths_offsets[commit_start]
379 }
380
381 #[inline]
382 pub(in crate::whir) fn num_commits(&self, proof_idx: usize) -> usize {
383 self.commits_per_proof_offsets[proof_idx + 1] - self.commits_per_proof_offsets[proof_idx]
384 }
385
386 #[inline]
388 pub(in crate::whir) fn decompose(&self, row_idx: usize) -> CodewordAccsIdx {
389 let proof_idx = self.rows_per_proof_offsets[1..].partition_point(|&x| x <= row_idx);
390 let record_idx = row_idx - self.rows_per_proof_offsets[proof_idx];
391 let commit_start = self.commits_per_proof_offsets[proof_idx];
392 let commit_end = self.commits_per_proof_offsets[proof_idx + 1];
393
394 let chunks_before_proof = self.stacking_chunks_offsets[commit_start];
395 let chunks_after_proof = self.stacking_chunks_offsets[commit_end];
396 let total_chunks_for_proof = chunks_after_proof - chunks_before_proof;
397
398 let query_idx = record_idx / (self.num_cosets * total_chunks_for_proof);
399 let coset_idx = (record_idx / total_chunks_for_proof) % self.num_cosets;
400 let local_chunk_idx = record_idx % total_chunks_for_proof;
401 let absolute_chunk_idx = chunks_before_proof + local_chunk_idx;
402 let commit_idx = self.stacking_chunks_offsets[commit_start + 1..=commit_end]
403 .partition_point(|&x| x <= absolute_chunk_idx);
404 let chunk_idx =
405 absolute_chunk_idx - self.stacking_chunks_offsets[commit_start + commit_idx];
406
407 (proof_idx, query_idx, coset_idx, commit_idx, chunk_idx)
408 }
409
410 #[inline]
412 pub(in crate::whir) fn commit_num_chunks(&self, proof_idx: usize, commit_idx: usize) -> usize {
413 let global_idx = self.commits_per_proof_offsets[proof_idx] + commit_idx;
414 self.stacking_chunks_offsets[global_idx + 1] - self.stacking_chunks_offsets[global_idx]
415 }
416
417 #[inline]
419 pub(in crate::whir) fn commit_width(&self, proof_idx: usize, commit_idx: usize) -> usize {
420 let global_idx = self.commits_per_proof_offsets[proof_idx] + commit_idx;
421 self.stacking_widths_offsets[global_idx + 1] - self.stacking_widths_offsets[global_idx]
422 }
423
424 #[inline]
426 pub(in crate::whir) fn commit_width_offset(
427 &self,
428 proof_idx: usize,
429 commit_idx: usize,
430 ) -> usize {
431 let commit_start = self.commits_per_proof_offsets[proof_idx];
432 self.stacking_widths_offsets[commit_start + commit_idx]
433 - self.stacking_widths_offsets[commit_start]
434 }
435
436 #[inline]
439 pub(in crate::whir) fn chunk_len(
440 &self,
441 proof_idx: usize,
442 commit_idx: usize,
443 chunk_idx: usize,
444 ) -> usize {
445 (self.commit_width(proof_idx, commit_idx) - chunk_idx * CHUNK).min(CHUNK)
446 }
447}
448
449impl FlattenedLayout for CodewordAccsLayout {
450 type Index = CodewordAccsIdx;
451
452 #[inline]
453 fn len(&self) -> usize {
454 *self.rows_per_proof_offsets.last().unwrap_or(&0)
455 }
456
457 #[inline]
458 fn offset(&self, (proof, query, coset, commit, chunk): Self::Index) -> usize {
459 debug_assert!(proof < self.num_proofs());
460 debug_assert!(query < self.num_queries);
461 debug_assert!(coset < self.num_cosets);
462 debug_assert!(commit < self.num_commits(proof));
463 debug_assert!(chunk < self.commit_num_chunks(proof, commit));
464
465 let commit_start = self.commits_per_proof_offsets[proof];
466 let chunks_before_proof = self.stacking_chunks_offsets[commit_start];
467 let total_chunks_for_proof = self.total_chunks(proof);
468 let local_commit_chunk_offset =
469 self.stacking_chunks_offsets[commit_start + commit] - chunks_before_proof;
470
471 self.rows_per_proof_offsets[proof]
472 + query * self.num_cosets * total_chunks_for_proof
473 + coset * total_chunks_for_proof
474 + local_commit_chunk_offset
475 + chunk
476 }
477}
478
479#[cfg(feature = "cuda")]
480mod cuda_abi;
481
482pub struct WhirModule {
483 params: SystemParams,
484 bus_inventory: BusInventory,
485
486 sumcheck_bus: WhirSumcheckBus,
488 verify_queries_bus: VerifyQueriesBus,
489 verify_query_bus: VerifyQueryBus,
490 folding_bus: WhirFoldingBus,
491 final_poly_mle_eval_bus: FinalPolyMleEvalBus,
492 final_poly_query_eval_bus: FinalPolyQueryEvalBus,
493
494 alpha_bus: WhirAlphaBus,
496 gamma_bus: WhirGammaBus,
497 query_bus: WhirQueryBus,
498 eq_alpha_u_bus: WhirEqAlphaUBus,
499 final_poly_bus: WhirFinalPolyBus,
500 final_poly_folding_bus: FinalPolyFoldingBus,
501}
502
503impl WhirModule {
504 pub fn new(
505 child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
506 b: &mut BusIndexManager,
507 bus_inventory: BusInventory,
508 ) -> Self {
509 let sumcheck_bus = WhirSumcheckBus::new(b.new_bus_idx());
510 let alpha_bus = WhirAlphaBus::new(b.new_bus_idx());
511 let gamma_bus = WhirGammaBus::new(b.new_bus_idx());
512 let query_bus = WhirQueryBus::new(b.new_bus_idx());
513 let verify_queries_bus = VerifyQueriesBus::new(b.new_bus_idx());
514 let verify_query_bus = VerifyQueryBus::new(b.new_bus_idx());
515 let eq_alpha_u_bus = WhirEqAlphaUBus::new(b.new_bus_idx());
516 let folding_bus = WhirFoldingBus::new(b.new_bus_idx());
517 let final_poly_mle_eval_bus = FinalPolyMleEvalBus::new(b.new_bus_idx());
518 let final_poly_query_eval_bus = FinalPolyQueryEvalBus::new(b.new_bus_idx());
519 let final_poly_bus = WhirFinalPolyBus::new(b.new_bus_idx());
520 let final_poly_folding_bus = FinalPolyFoldingBus::new(b.new_bus_idx());
521 Self {
522 params: child_vk.inner.params.clone(),
523 bus_inventory,
524 sumcheck_bus,
525 verify_queries_bus,
526 verify_query_bus,
527 final_poly_mle_eval_bus,
528 final_poly_query_eval_bus,
529 folding_bus,
530 alpha_bus,
531 gamma_bus,
532 query_bus,
533 eq_alpha_u_bus,
534 final_poly_bus,
535 final_poly_folding_bus,
536 }
537 }
538}
539
540impl WhirModule {
541 #[tracing::instrument(level = "trace", skip_all)]
542 pub fn run_preflight<TS: FiatShamirTranscript<BabyBearPoseidon2Config> + TranscriptHistory>(
543 &self,
544 proof: &Proof<BabyBearPoseidon2Config>,
545 preflight: &mut Preflight,
546 ts: &mut TS,
547 ) {
548 let WhirProof {
549 mu_pow_witness: _, whir_sumcheck_polys,
551 codeword_commits,
552 ood_values,
553 initial_round_opened_rows: _,
554 initial_round_merkle_proofs: _,
555 codeword_opened_values: _,
556 codeword_merkle_proofs: _,
557 folding_pow_witnesses,
558 query_phase_pow_witnesses,
559 final_poly,
560 } = &proof.whir_proof;
561
562 let k_whir = self.params.k_whir();
563 let num_queries_per_round = num_queries_per_round(&self.params);
564 let query_layout = WhirQueryLayout::new(1, &num_queries_per_round);
565 let num_whir_rounds = self.params.num_whir_rounds();
566 let total_queries = query_layout.queries_per_proof();
567 let mut gammas = Vec::with_capacity(num_whir_rounds);
568 let mut z0s = Vec::with_capacity(num_whir_rounds - 1);
569 let mut alphas = Vec::with_capacity(num_whir_rounds * k_whir);
570 let mut folding_pow_samples = Vec::with_capacity(num_whir_rounds * k_whir);
571 let mut query_pow_samples = Vec::with_capacity(num_whir_rounds);
572 let mut queries = Vec::with_capacity(total_queries);
573 let mut whir_round_tidx_per_round = Vec::with_capacity(num_whir_rounds);
574 let mut query_tidx_per_round = Vec::with_capacity(num_whir_rounds);
575
576 debug_assert_eq!(whir_sumcheck_polys.len(), num_whir_rounds * k_whir);
577 debug_assert_eq!(folding_pow_witnesses.len(), num_whir_rounds * k_whir);
578 debug_assert_eq!(query_phase_pow_witnesses.len(), num_whir_rounds);
579 debug_assert_eq!(ood_values.len(), num_whir_rounds - 1);
580 debug_assert_eq!(codeword_commits.len(), num_whir_rounds - 1);
581
582 for i in 0..num_whir_rounds {
583 whir_round_tidx_per_round.push(ts.len());
584
585 let num_round_queries = num_queries_per_round[i];
586 for j in 0..k_whir {
587 let evals = &whir_sumcheck_polys[i * k_whir + j];
588 let &[ev1, ev2] = evals;
589 ts.observe_ext(ev1);
590 ts.observe_ext(ev2);
591
592 folding_pow_samples.push(pow_observe_sample(
593 ts,
594 self.params.whir.folding_pow_bits,
595 folding_pow_witnesses[i * k_whir + j],
596 ));
597 alphas.push(ts.sample_ext());
598 }
599
600 if i != num_whir_rounds - 1 {
601 ts.observe_commit(codeword_commits[i]);
602 z0s.push(ts.sample_ext());
603 ts.observe_ext(ood_values[i]);
604 } else {
605 for coeff in final_poly {
606 ts.observe_ext(*coeff);
607 }
608 };
609
610 query_pow_samples.push(pow_observe_sample(
611 ts,
612 self.params.whir.query_phase_pow_bits,
613 query_phase_pow_witnesses[i],
614 ));
615 query_tidx_per_round.push(ts.len());
616
617 for _ in 0..num_round_queries {
618 queries.push(ts.sample());
619 }
620 gammas.push(ts.sample_ext());
621 }
622 preflight.whir = WhirPreflight {
623 whir_round_tidx_per_round,
624 query_tidx_per_round,
625 alphas,
626 z0s,
627 gammas,
628 folding_pow_samples,
629 query_pow_samples,
630 queries,
631 };
632 }
633}
634
635pub(crate) struct WhirBlobCpu {
636 whir_round_tidx_per_round: FlattenedVec<usize, PerProofLayout>,
638 query_tidx_per_round: FlattenedVec<usize, PerProofLayout>,
639 initial_claim_per_round: FlattenedVec<EF, PerProofLayout>,
640 post_sumcheck_claims: FlattenedVec<EF, PerProofLayout>,
641 pre_query_claims: FlattenedVec<EF, PerProofLayout>,
642 eq_partials: FlattenedVec<EF, PerProofLayout>,
643 final_poly_at_u: Vec<EF>,
644 zi_roots: FlattenedVec<F, WhirQueryLayout>,
645 zis: FlattenedVec<F, WhirQueryLayout>,
646 yis: FlattenedVec<EF, WhirQueryLayout>,
647 fold_records: FlattenedVec<FoldRecord, PerProofLayout>,
648
649 codeword_value_accs: FlattenedVec<EF, CodewordAccsLayout>,
651 mu_pows: FlattenedVec<EF, VariablePerProofLayout>,
653}
654
655struct WhirBlobBuilder {
656 whir_round_tidx_per_round: Vec<usize>,
657 query_tidx_per_round: Vec<usize>,
658 initial_claim_per_round: Vec<EF>,
659 post_sumcheck_claims: Vec<EF>,
660 pre_query_claims: Vec<EF>,
661 eq_partials: Vec<EF>,
662 final_poly_at_u: Vec<EF>,
663 zi_roots: Vec<F>,
664 zis: Vec<F>,
665 yis: Vec<EF>,
666 fold_records: Vec<FoldRecord>,
667 codeword_value_accs: Vec<EF>,
668 mu_pows: Vec<EF>,
669}
670
671struct WhirBlobLayouts {
672 tidx_layout: PerProofLayout,
673 initial_claim_layout: PerProofLayout,
674 pre_query_claim_layout: PerProofLayout,
675 sumcheck_layout: PerProofLayout,
676 query_layout: WhirQueryLayout,
677 fold_layout: PerProofLayout,
678 accs_layout: CodewordAccsLayout,
679 mu_pows_layout: VariablePerProofLayout,
680}
681
682impl WhirBlobBuilder {
683 fn with_capacities(
684 num_proofs: usize,
685 num_whir_rounds: usize,
686 sumcheck_rows_per_proof: usize,
687 queries_per_proof: usize,
688 fold_records_per_proof: usize,
689 codeword_value_accs_len: usize,
690 mu_pows_len: usize,
691 ) -> Self {
692 Self {
693 whir_round_tidx_per_round: Vec::with_capacity(num_proofs * num_whir_rounds),
694 query_tidx_per_round: Vec::with_capacity(num_proofs * num_whir_rounds),
695 initial_claim_per_round: Vec::with_capacity(num_proofs * (num_whir_rounds + 1)),
696 post_sumcheck_claims: Vec::with_capacity(num_proofs * sumcheck_rows_per_proof),
697 pre_query_claims: Vec::with_capacity(num_proofs * num_whir_rounds),
698 eq_partials: Vec::with_capacity(num_proofs * sumcheck_rows_per_proof),
699 final_poly_at_u: Vec::with_capacity(num_proofs),
700 zi_roots: Vec::with_capacity(num_proofs * queries_per_proof),
701 zis: Vec::with_capacity(num_proofs * queries_per_proof),
702 yis: Vec::with_capacity(num_proofs * queries_per_proof),
703 fold_records: Vec::with_capacity(num_proofs * fold_records_per_proof),
704 codeword_value_accs: Vec::with_capacity(codeword_value_accs_len),
705 mu_pows: Vec::with_capacity(mu_pows_len),
706 }
707 }
708
709 fn into_blob(self, layouts: WhirBlobLayouts) -> WhirBlobCpu {
710 let WhirBlobLayouts {
711 tidx_layout,
712 initial_claim_layout,
713 pre_query_claim_layout,
714 sumcheck_layout,
715 query_layout,
716 fold_layout,
717 accs_layout,
718 mu_pows_layout,
719 } = layouts;
720 WhirBlobCpu {
721 whir_round_tidx_per_round: FlattenedVec::from_parts(
722 tidx_layout.clone(),
723 self.whir_round_tidx_per_round,
724 ),
725 query_tidx_per_round: FlattenedVec::from_parts(tidx_layout, self.query_tidx_per_round),
726 initial_claim_per_round: FlattenedVec::from_parts(
727 initial_claim_layout,
728 self.initial_claim_per_round,
729 ),
730 post_sumcheck_claims: FlattenedVec::from_parts(
731 sumcheck_layout.clone(),
732 self.post_sumcheck_claims,
733 ),
734 pre_query_claims: FlattenedVec::from_parts(
735 pre_query_claim_layout,
736 self.pre_query_claims,
737 ),
738 eq_partials: FlattenedVec::from_parts(sumcheck_layout, self.eq_partials),
739 final_poly_at_u: self.final_poly_at_u,
740 zi_roots: FlattenedVec::from_parts(query_layout.clone(), self.zi_roots),
741 zis: FlattenedVec::from_parts(query_layout.clone(), self.zis),
742 yis: FlattenedVec::from_parts(query_layout, self.yis),
743 fold_records: FlattenedVec::from_parts(fold_layout, self.fold_records),
744 codeword_value_accs: FlattenedVec::from_parts(accs_layout, self.codeword_value_accs),
745 mu_pows: FlattenedVec::from_parts(mu_pows_layout, self.mu_pows),
746 }
747 }
748}
749
750impl AirModule for WhirModule {
751 fn num_airs(&self) -> usize {
752 WhirModuleChipDiscriminants::COUNT
753 }
754
755 fn airs<SC: StarkProtocolConfig<F = F>>(&self) -> Vec<AirRef<SC>> {
756 let params = &self.params;
757 let initial_log_domain_size = params.n_stack + params.l_skip + params.log_blowup;
758
759 let num_rounds = params.num_whir_rounds();
760 let num_queries_per_round = num_queries_per_round(params);
761
762 let whir_round_air: AirRef<SC> = Arc::new(WhirRoundAir {
763 whir_module_bus: self.bus_inventory.whir_module_bus,
764 commitments_bus: self.bus_inventory.commitments_bus,
765 transcript_bus: self.bus_inventory.transcript_bus,
766 exp_bits_len_bus: self.bus_inventory.exp_bits_len_bus,
767 sumcheck_bus: self.sumcheck_bus,
768 verify_queries_bus: self.verify_queries_bus,
769 final_poly_mle_eval_bus: self.final_poly_mle_eval_bus,
770 final_poly_query_eval_bus: self.final_poly_query_eval_bus,
771 query_bus: self.query_bus,
772 gamma_bus: self.gamma_bus,
773 k: params.k_whir(),
774 num_rounds,
775 initial_log_domain_size,
776 final_poly_len: 1 << params.log_final_poly_len(),
777 pow_bits: params.whir.query_phase_pow_bits,
778 folding_pow_bits: params.whir.folding_pow_bits,
779 generator: F::GENERATOR,
780 whir_round_encoder: whir_round_encoder(num_rounds),
781 num_queries_per_round: num_queries_per_round.clone(),
782 });
783 let whir_sumcheck_air = SumcheckAir {
784 sumcheck_bus: self.sumcheck_bus,
785 whir_opening_point_bus: self.bus_inventory.whir_opening_point_bus,
786 transcript_bus: self.bus_inventory.transcript_bus,
787 exp_bits_len_bus: self.bus_inventory.exp_bits_len_bus,
788 alpha_bus: self.alpha_bus,
789 eq_alpha_u_bus: self.eq_alpha_u_bus,
790 k: params.k_whir(),
791 folding_pow_bits: params.whir.folding_pow_bits,
792 generator: F::GENERATOR,
793 };
794 let initial_round_opened_values_air = InitialOpenedValuesAir {
795 stacking_indices_bus: self.bus_inventory.stacking_indices_bus,
796 whir_mu_bus: self.bus_inventory.whir_mu_bus,
797 verify_query_bus: self.verify_query_bus,
798 folding_bus: self.folding_bus,
799 poseidon_permute_bus: self.bus_inventory.poseidon2_permute_bus,
800 merkle_verify_bus: self.bus_inventory.merkle_verify_bus,
801 k: params.k_whir(),
802 initial_log_domain_size,
803 };
804 let non_initial_round_opened_values_air = NonInitialOpenedValuesAir {
805 verify_query_bus: self.verify_query_bus,
806 folding_bus: self.folding_bus,
807 poseidon2_compress_bus: self.bus_inventory.poseidon2_compress_bus,
808 merkle_verify_bus: self.bus_inventory.merkle_verify_bus,
809 k: params.k_whir(),
810 initial_log_domain_size,
811 };
812 let query_air = WhirQueryAir {
813 transcript_bus: self.bus_inventory.transcript_bus,
814 exp_bits_len_bus: self.bus_inventory.exp_bits_len_bus,
815 query_bus: self.query_bus,
816 verify_queries_bus: self.verify_queries_bus,
817 verify_query_bus: self.verify_query_bus,
818 k: params.k_whir(),
819 initial_log_domain_size,
820 };
821 let folding_air = WhirFoldingAir {
822 alpha_bus: self.alpha_bus,
823 folding_bus: self.folding_bus,
824 k: params.k_whir(),
825 };
826 let final_poly_mle_eval_air = FinalPolyMleEvalAir {
827 whir_opening_point_bus: self.bus_inventory.whir_opening_point_bus,
828 whir_opening_point_lookup_bus: self.bus_inventory.whir_opening_point_lookup_bus,
829 transcript_bus: self.bus_inventory.transcript_bus,
830 final_poly_mle_eval_bus: self.final_poly_mle_eval_bus,
831 eq_alpha_u_bus: self.eq_alpha_u_bus,
832 final_poly_bus: self.final_poly_bus,
833 folding_bus: self.final_poly_folding_bus,
834 num_vars: params.log_final_poly_len(),
835 num_sumcheck_rounds: params.num_whir_sumcheck_rounds(),
836 num_whir_rounds: params.num_whir_rounds(),
837 total_whir_queries: params
838 .whir
839 .rounds
840 .iter()
841 .map(|cfg| cfg.num_queries + 1)
842 .sum(),
843 };
844 let final_poly_query_eval_air = FinalPolyQueryEvalAir {
845 query_bus: self.query_bus,
846 alpha_bus: self.alpha_bus,
847 gamma_bus: self.gamma_bus,
848 final_poly_bus: self.final_poly_bus,
849 final_poly_query_eval_bus: self.final_poly_query_eval_bus,
850 num_whir_rounds: params.num_whir_rounds(),
851 k_whir: params.k_whir(),
852 log_final_poly_len: params.log_final_poly_len(),
853 };
854 vec![
855 whir_round_air,
856 Arc::new(whir_sumcheck_air),
857 Arc::new(query_air),
858 Arc::new(initial_round_opened_values_air),
859 Arc::new(non_initial_round_opened_values_air),
860 Arc::new(folding_air),
861 Arc::new(final_poly_mle_eval_air),
862 Arc::new(final_poly_query_eval_air),
863 ]
864 }
865}
866
867impl WhirModule {
868 fn append_derived_whir_data_for_proof(
869 proof: &Proof<BabyBearPoseidon2Config>,
870 preflight: &Preflight,
871 local_mu_pows: &[EF],
872 params: &SystemParams,
873 query_layout: &WhirQueryLayout,
874 blob: &mut WhirBlobBuilder,
875 ) {
876 let k_whir = params.k_whir();
877 let num_whir_rounds = params.num_whir_rounds();
878 let l_skip = params.l_skip;
879 let initial_log_rs_domain_size = params.l_skip + params.n_stack + params.log_blowup;
880
881 let mut sumcheck_poly_iter = proof.whir_proof.whir_sumcheck_polys.iter();
882 let mut claim = proof
883 .stacking_proof
884 .stacking_openings
885 .iter()
886 .flatten()
887 .zip(local_mu_pows.iter())
888 .fold(EF::ZERO, |acc, (&opening, &mu_pow)| acc + mu_pow * opening);
889
890 let u = preflight.stacking.sumcheck_rnd[0]
891 .exp_powers_of_2()
892 .take(l_skip)
893 .chain(preflight.stacking.sumcheck_rnd[1..].iter().copied())
894 .collect_vec();
895
896 let mut eq_partial = EF::ONE;
897 for (i, query_range) in query_layout.iter_round_query_ranges() {
898 let round_queries = &preflight.whir.queries[query_range];
899 let log_rs_domain_size = initial_log_rs_domain_size - i;
900
901 blob.initial_claim_per_round.push(claim);
902
903 for j in 0..k_whir {
904 let evals = sumcheck_poly_iter.next().unwrap();
905 let &[ev1, ev2] = evals;
906 let ev0 = claim - ev1;
907 let alpha = preflight.whir.alphas[i * k_whir + j];
908 let uj = u[i * k_whir + j];
909 eq_partial *= EF::ONE - alpha - uj.double() + EF::from_u8(3) * uj * alpha;
912 blob.eq_partials.push(eq_partial);
913
914 claim = interpolate_quadratic_at_012(&[ev0, ev1, ev2], alpha);
915 blob.post_sumcheck_claims.push(claim);
916 }
917
918 let gamma = preflight.whir.gammas[i];
919 if let Some(&y0) = proof.whir_proof.ood_values.get(i) {
920 claim += gamma * y0;
921 }
922 let mut gamma_pows = gamma.powers().skip(2);
923
924 blob.pre_query_claims.push(claim);
925
926 let omega = F::two_adic_generator(log_rs_domain_size);
927 let round_alphas = &preflight.whir.alphas[i * k_whir..(i + 1) * k_whir];
928 for (query_idx, &sample) in round_queries.iter().enumerate() {
929 let index = sample.as_canonical_u32() & ((1 << (log_rs_domain_size - k_whir)) - 1);
930 let zi_root = omega.exp_u64(index as u64);
931 let zi = zi_root.exp_power_of_2(k_whir);
932 let record_start = blob.fold_records.len();
933 let yi = if i == 0 {
934 let mut codeword_vals = vec![EF::ZERO; 1 << k_whir];
935 let mut mu_pow_iter = local_mu_pows.iter();
936 for opened_rows_per_query in proof.whir_proof.initial_round_opened_rows.iter() {
937 let opened_rows = &opened_rows_per_query[query_idx];
938 let width = opened_rows[0].len();
939 for c in 0..width {
940 let mu_pow = mu_pow_iter.next().unwrap();
941 for (cv, row) in codeword_vals.iter_mut().zip(opened_rows.iter()) {
942 *cv += *mu_pow * row[c];
943 }
944 }
945 }
946 binary_k_fold(
947 codeword_vals,
948 round_alphas,
949 zi_root,
950 i,
951 query_idx,
952 &mut blob.fold_records,
953 )
954 } else {
955 let opened_values =
956 proof.whir_proof.codeword_opened_values[i - 1][query_idx].clone();
957 binary_k_fold(
958 opened_values,
959 round_alphas,
960 zi_root,
961 i,
962 query_idx,
963 &mut blob.fold_records,
964 )
965 };
966 for rec in &mut blob.fold_records[record_start..] {
967 rec.set_final_values(zi, yi);
968 }
969 blob.zi_roots.push(zi_root);
970 blob.zis.push(zi);
971 blob.yis.push(yi);
972
973 claim += gamma_pows.next().unwrap() * yi;
974 }
975 let _ = gamma_pows.next().unwrap();
976 }
977
978 blob.initial_claim_per_round.push(claim);
980 debug_assert!(sumcheck_poly_iter.next().is_none());
981
982 let t = k_whir * num_whir_rounds;
985 blob.final_poly_at_u
986 .push(eval_final_poly_at_u(&proof.whir_proof.final_poly, &u[t..]));
987 }
988
989 fn enqueue_pow_requests_for_proof(
990 exp_bits_len_gen: &ExpBitsLenCpuTraceGenerator,
991 preflight: &Preflight,
992 params: &SystemParams,
993 query_layout: &WhirQueryLayout,
994 ) {
995 let mu_pow_bits = params.whir.mu_pow_bits;
996 let folding_pow_bits = params.whir.folding_pow_bits;
997 let query_phase_pow_bits = params.whir.query_phase_pow_bits;
998 let k_whir = params.k_whir();
999 let initial_log_rs_domain_size = params.l_skip + params.n_stack + params.log_blowup;
1000
1001 if mu_pow_bits > 0 {
1003 exp_bits_len_gen.add_requests(std::iter::once((
1004 F::GENERATOR,
1005 preflight.stacking.mu_pow_sample,
1006 mu_pow_bits,
1007 )));
1008 }
1009
1010 if folding_pow_bits > 0 {
1011 exp_bits_len_gen.add_requests(
1012 preflight
1013 .whir
1014 .folding_pow_samples
1015 .iter()
1016 .map(|pow_sample| (F::GENERATOR, *pow_sample, folding_pow_bits)),
1017 );
1018 }
1019 if query_phase_pow_bits > 0 {
1020 exp_bits_len_gen.add_requests(
1021 preflight
1022 .whir
1023 .query_pow_samples
1024 .iter()
1025 .map(|pow_sample| (F::GENERATOR, *pow_sample, query_phase_pow_bits)),
1026 );
1027 }
1028
1029 for (i, query_range) in query_layout.iter_round_query_ranges() {
1030 let round_queries = &preflight.whir.queries[query_range];
1031 let log_rs_domain_size = initial_log_rs_domain_size - i;
1032 let omega = F::two_adic_generator(log_rs_domain_size);
1033 let shift_bits = initial_log_rs_domain_size - k_whir + 1 - i;
1034 let per_query_lookups = if i == 0 {
1035 preflight.initial_row_states.len() as u32
1036 } else {
1037 1
1038 };
1039 exp_bits_len_gen.add_requests_with_shift(round_queries.iter().copied().map(|sample| {
1040 (
1041 omega,
1042 sample,
1043 log_rs_domain_size - k_whir,
1044 shift_bits,
1045 per_query_lookups,
1046 )
1047 }));
1048 }
1049 }
1050
1051 fn append_initial_opened_values_accs_for_proof(
1052 blob: &mut WhirBlobBuilder,
1053 accs_layout: &CodewordAccsLayout,
1054 proof_idx: usize,
1055 proof: &Proof<BabyBearPoseidon2Config>,
1056 local_mu_pows: &[EF],
1057 num_initial_queries: usize,
1058 k_whir: usize,
1059 ) {
1060 debug_assert_eq!(
1061 proof.whir_proof.initial_round_opened_rows.len(),
1062 accs_layout.num_commits(proof_idx)
1063 );
1064 #[cfg(debug_assertions)]
1065 for (i, openings_per_commit) in proof
1066 .whir_proof
1067 .initial_round_opened_rows
1068 .iter()
1069 .enumerate()
1070 {
1071 debug_assert_eq!(
1072 openings_per_commit[0][0].len(),
1073 accs_layout.commit_width(proof_idx, i)
1074 );
1075 }
1076
1077 for query_idx in 0..num_initial_queries {
1078 let mut codeword_vals = EF::zero_vec(1 << k_whir);
1079 for (coset_idx, codeword_val) in codeword_vals.iter_mut().enumerate() {
1080 let mut base = 0;
1081 for opened_rows_per_query in proof.whir_proof.initial_round_opened_rows.iter() {
1082 let opened_rows = &opened_rows_per_query[query_idx];
1083 let width = opened_rows[0].len();
1084 let num_chunks = width.div_ceil(CHUNK);
1085
1086 for chunk_idx in 0..num_chunks {
1087 let chunk_start = chunk_idx * CHUNK;
1088 let chunk_len = cmp::min(CHUNK, width - chunk_start);
1089
1090 let opened_chunk =
1091 &opened_rows[coset_idx][chunk_start..chunk_start + chunk_len];
1092
1093 blob.codeword_value_accs.push(*codeword_val);
1094
1095 for (offset, &val) in opened_chunk.iter().enumerate() {
1096 *codeword_val += local_mu_pows[base + chunk_start + offset] * val;
1097 }
1098 }
1099 base += width;
1100 }
1101 }
1102 }
1103 }
1104
1105 #[tracing::instrument(skip_all)]
1106 fn generate_blob(
1107 &self,
1108 child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
1109 proofs: &[&Proof<BabyBearPoseidon2Config>],
1110 preflights: &[&Preflight],
1111 exp_bits_len_gen: &ExpBitsLenCpuTraceGenerator,
1112 ) -> WhirBlobCpu {
1113 let params = &child_vk.inner.params;
1114 let k_whir = params.k_whir();
1115 let num_queries_per_round = num_queries_per_round(params);
1116 let num_initial_queries = *num_queries_per_round.first().unwrap_or(&0);
1117 let num_whir_rounds = params.num_whir_rounds();
1118 let query_layout = WhirQueryLayout::new(proofs.len(), &num_queries_per_round);
1119 let queries_per_proof = query_layout.queries_per_proof();
1120 let sumcheck_rows_per_proof = params.num_whir_sumcheck_rounds();
1121 let fold_records_per_proof = queries_per_proof * ((1 << k_whir) - 1);
1122 let tidx_layout = PerProofLayout::new(proofs.len(), num_whir_rounds);
1123
1124 let per_proof_commit_widths: Vec<Vec<usize>> = proofs
1125 .iter()
1126 .map(|proof| {
1127 proof
1128 .whir_proof
1129 .initial_round_opened_rows
1130 .iter()
1131 .map(|openings| openings[0][0].len())
1132 .collect()
1133 })
1134 .collect();
1135 let accs_layout =
1136 CodewordAccsLayout::new(num_initial_queries, 1 << k_whir, &per_proof_commit_widths);
1137 let mu_pows_layout =
1138 VariablePerProofLayout::new((0..proofs.len()).map(|i| accs_layout.total_width(i)));
1139
1140 let mut blob = WhirBlobBuilder::with_capacities(
1141 proofs.len(),
1142 num_whir_rounds,
1143 sumcheck_rows_per_proof,
1144 queries_per_proof,
1145 fold_records_per_proof,
1146 accs_layout.len(),
1147 mu_pows_layout.len(),
1148 );
1149
1150 for (proof_idx, (proof, preflight)) in zip(proofs, preflights).enumerate() {
1151 let mu = preflight.stacking.stacking_batching_challenge;
1152 let local_mu_pows = mu
1153 .powers()
1154 .take(accs_layout.total_width(proof_idx))
1155 .collect_vec();
1156
1157 Self::append_derived_whir_data_for_proof(
1158 proof,
1159 preflight,
1160 &local_mu_pows,
1161 params,
1162 &query_layout,
1163 &mut blob,
1164 );
1165
1166 blob.whir_round_tidx_per_round
1167 .extend_from_slice(&preflight.whir.whir_round_tidx_per_round);
1168 blob.query_tidx_per_round
1169 .extend_from_slice(&preflight.whir.query_tidx_per_round);
1170
1171 Self::enqueue_pow_requests_for_proof(
1172 exp_bits_len_gen,
1173 preflight,
1174 params,
1175 &query_layout,
1176 );
1177
1178 Self::append_initial_opened_values_accs_for_proof(
1179 &mut blob,
1180 &accs_layout,
1181 proof_idx,
1182 proof,
1183 &local_mu_pows,
1184 num_initial_queries,
1185 k_whir,
1186 );
1187
1188 blob.mu_pows.extend(local_mu_pows);
1189 }
1190 let initial_claim_layout = PerProofLayout::new(proofs.len(), num_whir_rounds + 1);
1191 let pre_query_claim_layout = PerProofLayout::new(proofs.len(), num_whir_rounds);
1192 let sumcheck_layout = PerProofLayout::new(proofs.len(), sumcheck_rows_per_proof);
1193 let fold_layout = PerProofLayout::new(proofs.len(), fold_records_per_proof);
1194 let layouts = WhirBlobLayouts {
1195 tidx_layout,
1196 initial_claim_layout,
1197 pre_query_claim_layout,
1198 sumcheck_layout,
1199 query_layout,
1200 fold_layout,
1201 accs_layout,
1202 mu_pows_layout,
1203 };
1204
1205 blob.into_blob(layouts)
1206 }
1207}
1208
1209impl<SC: StarkProtocolConfig<F = F>> TraceGenModule<GlobalCtxCpu, CpuBackend<SC>> for WhirModule {
1210 type ModuleSpecificCtx<'a> = ExpBitsLenCpuTraceGenerator;
1211
1212 #[tracing::instrument(skip_all)]
1213 fn generate_proving_ctxs(
1214 &self,
1215 child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
1216 proofs: &[Proof<BabyBearPoseidon2Config>],
1217 preflights: &[Preflight],
1218 exp_bits_len_gen: &ExpBitsLenCpuTraceGenerator,
1219 required_heights: Option<&[usize]>,
1220 ) -> Option<Vec<AirProvingContext<CpuBackend<SC>>>> {
1221 let proofs = proofs.iter().collect_vec();
1222 let preflights = preflights.iter().collect_vec();
1223 let blob = self.generate_blob(child_vk, &proofs, &preflights, exp_bits_len_gen);
1224 let ctx = (
1225 StandardTracegenCtx {
1226 vk: child_vk,
1227 proofs: &proofs,
1228 preflights: &preflights,
1229 },
1230 &blob,
1231 );
1232
1233 let chips = [
1234 WhirModuleChip::WhirRound,
1235 WhirModuleChip::Sumcheck,
1236 WhirModuleChip::Query,
1237 WhirModuleChip::InitialOpenedValues,
1238 WhirModuleChip::NonInitialOpenedValues,
1239 WhirModuleChip::Folding,
1240 WhirModuleChip::FinalPolyMleEval,
1241 WhirModuleChip::FinalPolyQueryEval,
1242 ];
1243 let span = tracing::Span::current();
1244 chips
1245 .par_iter()
1246 .map(|chip| {
1247 let _guard = span.enter();
1248 chip.generate_proving_ctx(
1249 &ctx,
1250 required_heights.map(|heights| heights[chip.index()]),
1251 )
1252 })
1253 .collect::<Vec<_>>()
1254 .into_iter()
1255 .collect()
1256 }
1257}
1258
1259fn binary_k_fold(
1260 mut values: Vec<EF>,
1261 alphas: &[EF],
1262 base_coset_shift: F,
1263 whir_round: usize,
1264 query_idx: usize,
1265 records: &mut Vec<FoldRecord>,
1266) -> EF {
1267 let n = values.len();
1268 let k = alphas.len();
1269 debug_assert_eq!(n, 1 << k);
1270
1271 let omega_k = F::two_adic_generator(k);
1272 let omega_k_inv = omega_k.inverse();
1273
1274 let tw = omega_k.powers().take(1 << (k - 1)).collect_vec();
1275 let inv_tw = omega_k_inv.powers().take(1 << (k - 1)).collect_vec();
1276
1277 for (j, (&alpha, coset_shift, coset_shift_inv)) in izip!(
1278 alphas.iter(),
1279 base_coset_shift.exp_powers_of_2(),
1280 base_coset_shift.inverse().exp_powers_of_2()
1281 )
1282 .enumerate()
1283 {
1284 let m = n >> (j + 1);
1285 let (lo, hi) = values.split_at_mut(m);
1286
1287 for i in 0..m {
1288 let eval_point = tw[i << j] * coset_shift;
1289 let eval_point_inv = inv_tw[i << j] * coset_shift_inv;
1290 let new_val = lo[i] + (alpha - eval_point) * (lo[i] - hi[i]) * eval_point_inv.halve();
1291 records.push(FoldRecord::new(
1292 whir_round,
1293 query_idx,
1294 tw[i << j],
1295 coset_shift,
1296 m,
1297 i,
1298 j + 1,
1299 lo[i],
1300 hi[i],
1301 new_val,
1302 alpha,
1303 ));
1304 lo[i] = new_val;
1305 }
1306 }
1307 values[0]
1308}
1309
1310#[derive(Clone, Copy, strum_macros::Display, EnumDiscriminants)]
1311#[strum_discriminants(derive(strum_macros::EnumCount))]
1312#[strum_discriminants(repr(usize))]
1313enum WhirModuleChip {
1314 WhirRound,
1315 Sumcheck,
1316 Query,
1317 InitialOpenedValues,
1318 NonInitialOpenedValues,
1319 Folding,
1320 FinalPolyMleEval,
1321 FinalPolyQueryEval,
1322}
1323
1324impl WhirModuleChip {
1325 fn index(&self) -> usize {
1326 WhirModuleChipDiscriminants::from(self) as usize
1327 }
1328}
1329
1330impl RowMajorChip<F> for WhirModuleChip {
1331 type Ctx<'a> = (StandardTracegenCtx<'a>, &'a WhirBlobCpu);
1332
1333 #[tracing::instrument(
1334 name = "wrapper.generate_trace",
1335 level = "trace",
1336 skip_all,
1337 fields(air = %self)
1338 )]
1339 fn generate_trace(
1340 &self,
1341 ctx: &Self::Ctx<'_>,
1342 required_height: Option<usize>,
1343 ) -> Option<RowMajorMatrix<F>> {
1344 use WhirModuleChip::*;
1345 match self {
1346 WhirRound => whir_round::WhirRoundTraceGenerator.generate_trace(ctx, required_height),
1347 Sumcheck => sumcheck::WhirSumcheckTraceGenerator.generate_trace(ctx, required_height),
1348 Query => query::WhirQueryTraceGenerator.generate_trace(ctx, required_height),
1349 InitialOpenedValues => {
1350 let initial_opened_values_ctx = initial_opened_values::InitialOpenedValuesCtx {
1351 vk: ctx.0.vk,
1352 proofs: ctx.0.proofs,
1353 preflights: ctx.0.preflights,
1354 blob: ctx.1,
1355 };
1356 initial_opened_values::InitialOpenedValuesTraceGenerator
1357 .generate_trace(&initial_opened_values_ctx, required_height)
1358 }
1359 NonInitialOpenedValues => {
1360 non_initial_opened_values::NonInitialOpenedValuesTraceGenerator
1361 .generate_trace(ctx, required_height)
1362 }
1363 Folding => folding::FoldingTraceGenerator.generate_trace(ctx, required_height),
1364 FinalPolyMleEval => final_poly_mle_eval::FinalPolyMleEvalTraceGenerator
1365 .generate_trace(ctx, required_height),
1366 FinalPolyQueryEval => {
1367 let records = final_poly_query_eval::build_final_poly_query_eval_records(
1368 &ctx.0.vk.inner.params,
1369 ctx.0.proofs,
1370 ctx.0.preflights,
1371 &ctx.1.zis,
1372 );
1373 let final_poly_query_eval_ctx = final_poly_query_eval::FinalPolyQueryEvalCtx {
1374 vk: ctx.0.vk,
1375 preflights: ctx.0.preflights,
1376 records: records.as_slice(),
1377 };
1378 final_poly_query_eval::FinalPolyQueryEvalTraceGenerator
1379 .generate_trace(&final_poly_query_eval_ctx, required_height)
1380 }
1381 }
1382 }
1383}
1384
1385#[cfg(feature = "cuda")]
1386mod cuda_tracegen {
1387 use std::cmp;
1388
1389 use openvm_cuda_backend::{data_transporter::transport_matrix_h2d_row, GpuBackend};
1390 use openvm_cuda_common::{d_buffer::DeviceBuffer, stream::GpuDeviceCtx};
1391 use openvm_poseidon2_air::POSEIDON2_WIDTH;
1392 use openvm_stark_backend::p3_maybe_rayon::prelude::*;
1393 use openvm_stark_sdk::config::baby_bear_poseidon2::CHUNK;
1394
1395 use super::*;
1396 use crate::{
1397 cuda::{
1398 preflight::PreflightGpu, proof::ProofGpu, to_device_or_nullptr_on, vk::VerifyingKeyGpu,
1399 GlobalCtxGpu,
1400 },
1401 tracegen::{cuda::StandardTracegenGpuCtx, RowMajorChip, StandardTracegenCtx},
1402 whir::cuda_abi::PoseidonStatePair,
1403 };
1404
1405 pub(in crate::whir) struct WhirBlobGpu {
1406 pub zis: DeviceBuffer<F>,
1407 pub zi_roots: DeviceBuffer<F>,
1408 pub yis: DeviceBuffer<EF>,
1409 pub raw_queries: DeviceBuffer<F>,
1410 pub codeword_value_accs: DeviceBuffer<EF>,
1411 pub poseidon_states: DeviceBuffer<PoseidonStatePair>,
1412 pub rows_per_proof_offsets: DeviceBuffer<usize>,
1413 pub commits_per_proof_offsets: DeviceBuffer<usize>,
1414 pub stacking_chunks_offsets: DeviceBuffer<usize>,
1415 pub stacking_widths_offsets: DeviceBuffer<usize>,
1416 pub mus: DeviceBuffer<EF>,
1417 pub mu_pows: DeviceBuffer<EF>,
1418 pub codeword_opened_values: DeviceBuffer<EF>,
1419 pub codeword_states: DeviceBuffer<F>,
1420 pub folding_records: DeviceBuffer<FoldRecord>,
1421 }
1422
1423 fn build_initial_poseidon_state_pairs(
1424 proofs: &[&Proof<BabyBearPoseidon2Config>],
1425 preflights: &[&Preflight],
1426 ) -> Vec<PoseidonStatePair> {
1427 let expected_pairs: usize = preflights
1428 .iter()
1429 .flat_map(|preflight| {
1430 preflight
1431 .initial_row_states
1432 .iter()
1433 .flat_map(|commit| commit.iter().flat_map(|query| query.iter()))
1434 })
1435 .map(|coset_states| coset_states.len())
1436 .sum();
1437
1438 let mut pairs = Vec::with_capacity(expected_pairs);
1439 for (proof, preflight) in proofs.iter().zip(preflights) {
1440 if preflight.initial_row_states.is_empty() {
1441 continue;
1442 }
1443 let num_commits = preflight.initial_row_states.len();
1444 let num_queries = preflight.initial_row_states[0].len();
1445 let num_cosets = preflight.initial_row_states[0][0].len();
1446
1447 for query_idx in 0..num_queries {
1448 for coset_idx in 0..num_cosets {
1449 for commit_idx in 0..num_commits {
1450 let chunk_states =
1451 &preflight.initial_row_states[commit_idx][query_idx][coset_idx];
1452 let opened_row = &proof.whir_proof.initial_round_opened_rows[commit_idx]
1453 [query_idx][coset_idx];
1454 for (chunk_idx, &post_state) in chunk_states.iter().enumerate() {
1455 let mut pre_state = if chunk_idx > 0 {
1456 chunk_states[chunk_idx - 1]
1457 } else {
1458 [F::ZERO; POSEIDON2_WIDTH]
1459 };
1460 let chunk_start = chunk_idx * CHUNK;
1461 let chunk_len = cmp::min(CHUNK, opened_row.len() - chunk_start);
1462 pre_state[..chunk_len]
1463 .copy_from_slice(&opened_row[chunk_start..chunk_start + chunk_len]);
1464 pairs.push(PoseidonStatePair {
1465 pre_state,
1466 post_state,
1467 });
1468 }
1469 }
1470 }
1471 }
1472 }
1473 pairs
1474 }
1475
1476 impl WhirBlobGpu {
1477 fn new(
1478 proofs: &[&Proof<BabyBearPoseidon2Config>],
1479 preflights: &[&Preflight],
1480 blob: &WhirBlobCpu,
1481 device_ctx: &GpuDeviceCtx,
1482 ) -> Self {
1483 let mus = to_device_or_nullptr_on(
1484 &preflights
1485 .iter()
1486 .map(|preflight| preflight.stacking.stacking_batching_challenge)
1487 .collect_vec(),
1488 device_ctx,
1489 )
1490 .unwrap();
1491 let zis = to_device_or_nullptr_on(blob.zis.as_slice(), device_ctx).unwrap();
1492 let zi_roots = to_device_or_nullptr_on(blob.zi_roots.as_slice(), device_ctx).unwrap();
1493 let yis = to_device_or_nullptr_on(blob.yis.as_slice(), device_ctx).unwrap();
1494 let raw_queries = to_device_or_nullptr_on(
1495 &preflights
1496 .iter()
1497 .flat_map(|preflight| preflight.whir.queries.iter().copied())
1498 .collect_vec(),
1499 device_ctx,
1500 )
1501 .unwrap();
1502 let accs_layout = blob.codeword_value_accs.layout();
1503 let rows_per_proof_offsets =
1504 to_device_or_nullptr_on(accs_layout.rows_per_proof_offsets(), device_ctx).unwrap();
1505 let commits_per_proof_offsets =
1506 to_device_or_nullptr_on(accs_layout.commits_per_proof_offsets(), device_ctx)
1507 .unwrap();
1508 let stacking_chunks_offsets =
1509 to_device_or_nullptr_on(accs_layout.stacking_chunks_offsets(), device_ctx).unwrap();
1510 let stacking_widths_offsets =
1511 to_device_or_nullptr_on(accs_layout.stacking_widths_offsets(), device_ctx).unwrap();
1512 let mu_pows = to_device_or_nullptr_on(blob.mu_pows.as_slice(), device_ctx).unwrap();
1513 let codeword_value_accs =
1514 to_device_or_nullptr_on(blob.codeword_value_accs.as_slice(), device_ctx).unwrap();
1515
1516 let poseidon_states_host = build_initial_poseidon_state_pairs(proofs, preflights);
1518 let poseidon_states =
1519 to_device_or_nullptr_on(&poseidon_states_host, device_ctx).unwrap();
1520
1521 let folding_records =
1522 to_device_or_nullptr_on(blob.fold_records.as_slice(), device_ctx).unwrap();
1523
1524 let codeword_opened_values_cap: usize = proofs
1525 .iter()
1526 .map(|p| {
1527 p.whir_proof
1528 .codeword_opened_values
1529 .iter()
1530 .map(|r| r.iter().map(|q| q.len()).sum::<usize>())
1531 .sum::<usize>()
1532 })
1533 .sum();
1534 let mut codeword_opened_values_host = Vec::with_capacity(codeword_opened_values_cap);
1535 for proof in proofs.iter() {
1536 for round in proof.whir_proof.codeword_opened_values.iter() {
1537 for query in round.iter() {
1538 codeword_opened_values_host.extend_from_slice(query);
1539 }
1540 }
1541 }
1542 let codeword_opened_values =
1543 to_device_or_nullptr_on(&codeword_opened_values_host, device_ctx).unwrap();
1544
1545 let mut codeword_states_host =
1547 Vec::with_capacity(codeword_opened_values_cap * POSEIDON2_WIDTH);
1548 for preflight in preflights.iter() {
1549 for round in preflight.codeword_states.iter() {
1550 for query in round.iter() {
1551 for state in query.iter() {
1552 codeword_states_host.extend_from_slice(state);
1553 }
1554 }
1555 }
1556 }
1557 let codeword_states =
1558 to_device_or_nullptr_on(&codeword_states_host, device_ctx).unwrap();
1559
1560 WhirBlobGpu {
1561 zis,
1562 zi_roots,
1563 yis,
1564 raw_queries,
1565 codeword_value_accs,
1566 poseidon_states,
1567 rows_per_proof_offsets,
1568 commits_per_proof_offsets,
1569 stacking_chunks_offsets,
1570 stacking_widths_offsets,
1571 mus,
1572 mu_pows,
1573 folding_records,
1574 codeword_opened_values,
1575 codeword_states,
1576 }
1577 }
1578 }
1579
1580 impl ModuleChip<GpuBackend> for WhirModuleChip {
1581 type Ctx<'a> = (StandardTracegenGpuCtx<'a>, &'a WhirBlobGpu, &'a WhirBlobCpu);
1582
1583 fn generate_proving_ctx(
1584 &self,
1585 ctx: &Self::Ctx<'_>,
1586 required_height: Option<usize>,
1587 ) -> Option<AirProvingContext<GpuBackend>> {
1588 match self {
1589 WhirModuleChip::InitialOpenedValues => {
1590 initial_opened_values::cuda::InitialOpenedValuesGpuTraceGenerator
1591 .generate_proving_ctx(
1592 &initial_opened_values::cuda::InitialOpenedValuesGpuCtx {
1593 num_proofs: ctx.0.proofs.len(),
1594 blob: ctx.1,
1595 params: &ctx.0.vk.system_params,
1596 device_ctx: ctx.0.device_ctx,
1597 },
1598 required_height,
1599 )
1600 }
1601 WhirModuleChip::NonInitialOpenedValues => {
1602 non_initial_opened_values::cuda::NonInitialOpenedValuesGpuTraceGenerator
1603 .generate_proving_ctx(
1604 &non_initial_opened_values::cuda::NonInitialOpenedValuesGpuCtx {
1605 blob: ctx.1,
1606 params: &ctx.0.vk.system_params,
1607 device_ctx: ctx.0.device_ctx,
1608 },
1609 required_height,
1610 )
1611 }
1612 WhirModuleChip::FinalPolyQueryEval => {
1613 let proofs_cpu = ctx.0.proofs.iter().map(|proof| &proof.cpu).collect_vec();
1614 let preflights_cpu = ctx
1615 .0
1616 .preflights
1617 .iter()
1618 .map(|preflight| &preflight.cpu)
1619 .collect_vec();
1620 let records = final_poly_query_eval::build_final_poly_query_eval_records(
1621 &ctx.0.vk.system_params,
1622 &proofs_cpu,
1623 &preflights_cpu,
1624 &ctx.2.zis,
1625 );
1626 final_poly_query_eval::cuda::FinalPolyQueryEvalGpuTraceGenerator
1627 .generate_proving_ctx(
1628 &final_poly_query_eval::cuda::FinalPolyQueryEvalGpuCtx {
1629 records: records.as_slice(),
1630 params: &ctx.0.vk.system_params,
1631 preflights: ctx.0.preflights,
1632 device_ctx: ctx.0.device_ctx,
1633 },
1634 required_height,
1635 )
1636 }
1637 WhirModuleChip::Folding => folding::cuda::FoldingGpuTraceGenerator
1638 .generate_proving_ctx(
1639 &folding::cuda::FoldingGpuCtx {
1640 blob: ctx.1,
1641 params: &ctx.0.vk.system_params,
1642 num_proofs: ctx.0.proofs.len(),
1643 device_ctx: ctx.0.device_ctx,
1644 },
1645 required_height,
1646 ),
1647 _ => {
1648 let proofs_cpu = ctx.0.proofs.iter().map(|p| &p.cpu).collect_vec();
1649 let preflights_cpu = ctx.0.preflights.iter().map(|p| &p.cpu).collect_vec();
1650 let cpu_ctx = (
1651 StandardTracegenCtx {
1652 vk: &ctx.0.vk.cpu,
1653 proofs: &proofs_cpu,
1654 preflights: &preflights_cpu,
1655 },
1656 ctx.2,
1657 );
1658 let trace = RowMajorChip::generate_trace(self, &cpu_ctx, required_height);
1659 trace.map(|m| {
1660 AirProvingContext::simple_no_pis(
1661 transport_matrix_h2d_row(&m, ctx.0.device_ctx).unwrap(),
1662 )
1663 })
1664 }
1665 }
1666 }
1667 }
1668
1669 impl TraceGenModule<GlobalCtxGpu, GpuBackend> for WhirModule {
1670 type ModuleSpecificCtx<'a> = (
1671 &'a GpuExpBitsLenTraceGenerator,
1672 &'a openvm_cuda_common::stream::GpuDeviceCtx,
1673 );
1674
1675 #[tracing::instrument(skip_all)]
1676 fn generate_proving_ctxs(
1677 &self,
1678 child_vk: &VerifyingKeyGpu,
1679 proofs: &[ProofGpu],
1680 preflights: &[PreflightGpu],
1681 module_ctx: &Self::ModuleSpecificCtx<'_>,
1682 required_heights: Option<&[usize]>,
1683 ) -> Option<Vec<AirProvingContext<GpuBackend>>> {
1684 let exp_bits_len_gen = module_ctx.0;
1685 let device_ctx = module_ctx.1;
1686 let proofs_cpu = proofs.iter().map(|proof| &proof.cpu).collect_vec();
1687 let preflights_cpu = preflights
1688 .iter()
1689 .map(|preflight| &preflight.cpu)
1690 .collect_vec();
1691
1692 let blob = self.generate_blob(
1693 &child_vk.cpu,
1694 &proofs_cpu,
1695 &preflights_cpu,
1696 exp_bits_len_gen,
1697 );
1698 let blob_gpu = WhirBlobGpu::new(&proofs_cpu, &preflights_cpu, &blob, device_ctx);
1699 let ctx = (
1700 StandardTracegenGpuCtx {
1701 vk: child_vk,
1702 proofs,
1703 preflights,
1704 device_ctx,
1705 },
1706 &blob_gpu,
1707 &blob,
1708 );
1709
1710 let gpu_chips = [
1711 WhirModuleChip::InitialOpenedValues,
1712 WhirModuleChip::NonInitialOpenedValues,
1713 WhirModuleChip::Folding,
1714 WhirModuleChip::FinalPolyQueryEval,
1715 ];
1716 let cpu_chips = [
1717 WhirModuleChip::WhirRound,
1718 WhirModuleChip::Sumcheck,
1719 WhirModuleChip::Query,
1720 WhirModuleChip::FinalPolyMleEval,
1721 ];
1722
1723 let indexed_gpu_proving_ctxs = gpu_chips
1725 .iter()
1726 .map(|chip| {
1727 (
1728 chip.index(),
1729 chip.generate_proving_ctx(
1730 &ctx,
1731 required_heights.map(|heights| heights[chip.index()]),
1732 ),
1733 )
1734 })
1735 .collect::<Vec<_>>();
1736
1737 let cpu_proofs = ctx.0.proofs.iter().map(|p| &p.cpu).collect_vec();
1739 let cpu_preflights = ctx.0.preflights.iter().map(|p| &p.cpu).collect_vec();
1740 let cpu_ctx = (
1741 StandardTracegenCtx {
1742 vk: &ctx.0.vk.cpu,
1743 proofs: &cpu_proofs,
1744 preflights: &cpu_preflights,
1745 },
1746 ctx.2,
1747 );
1748 let span = tracing::Span::current();
1749 let indexed_cpu_rm_traces = cpu_chips
1750 .par_iter()
1751 .map(|chip| {
1752 let _guard = span.enter();
1753 (
1754 chip.index(),
1755 RowMajorChip::generate_trace(
1756 chip,
1757 &cpu_ctx,
1758 required_heights.map(|heights| heights[chip.index()]),
1759 ),
1760 )
1761 })
1762 .collect::<Vec<_>>();
1763
1764 let indexed_cpu_gpu_proving_ctxs = indexed_cpu_rm_traces
1766 .into_iter()
1767 .map(|(idx, trace)| {
1768 (
1769 idx,
1770 trace.map(|m| {
1771 AirProvingContext::simple_no_pis(
1772 transport_matrix_h2d_row(&m, device_ctx).unwrap(),
1773 )
1774 }),
1775 )
1776 })
1777 .collect::<Vec<_>>();
1778
1779 indexed_gpu_proving_ctxs
1780 .into_iter()
1781 .chain(indexed_cpu_gpu_proving_ctxs)
1782 .sorted_by(|a, b| a.0.cmp(&b.0))
1783 .map(|(_idx, ctx)| ctx)
1784 .collect()
1785 }
1786 }
1787}