openvm_recursion_circuit/whir/
mod.rs

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 requires at least 2 flags to work correctly.
68    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    /// Returns the raw query offset array (length = num_rounds + 1).
238    #[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/// Layout for the flattened `codeword_value_accs` array in `InitialOpenedValues`.
284/// Data is ordered as `[proof][query][coset][commit][chunk]`.
285#[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    /// Decompose a flat row index into `(proof, query, coset, commit, chunk)`.
387    #[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    /// Number of chunks in the given commit.
411    #[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    /// Width (number of opened values) in the given commit.
418    #[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    /// Cumulative width offset for the given commit (for `mu_pows` indexing).
425    #[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    /// Length of the chunk at `(commit_idx, chunk_idx)`.
437    /// The last chunk of a commit may be shorter than `CHUNK`.
438    #[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    // "execution" buses
487    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    // data buses
495    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: _, // Handled in stacking module preflight
550            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    // Flattened per-proof WHIR-derived data.
637    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    /// Initial opened values data with layout encoding (proof, query, coset, commit, chunk).
650    codeword_value_accs: FlattenedVec<EF, CodewordAccsLayout>,
651    /// Flattened as `[proof][mu_power_idx]` with per-proof variable widths.
652    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                // Möbius eq kernel: mobius_eq_1(u, alpha) = (1 - 2*u)*(1 - alpha) + u*alpha
910                //                              = 1 - alpha - 2*u + 3*u*alpha
911                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        // Push one for the final claim.
979        blob.initial_claim_per_round.push(claim);
980        debug_assert!(sumcheck_poly_iter.next().is_none());
981
982        // Evaluate the MLE of the table `final_poly` (interpreted as hypercube evaluations)
983        // at `u[t..]`. This matches the eval-to-coeff RS encoding semantics.
984        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        // μ PoW lookup (from stacking module)
1002        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            // Build poseidon state pairs in kernel order: [query][coset][commit][chunk]
1517            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            // Must be in same order as codeword_opened_values
1546            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            // Launch all CUDA tracegen kernels serially first (default stream).
1724            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            // Phase 1: CPU trace generation in parallel
1738            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            // Phase 2: H2D transfer serially on main thread
1765            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}