openvm_recursion_circuit/gkr/
mod.rs

1//! # GKR Air Module
2//!
3//! The GKR protocol reduces a fractional sum claim $\sum_{y \in H_{\ell+n}}
4//! \frac{\hat{p}(y)}{\hat{q}(y)} = 0$ to evaluation claims on the input layer polynomials at a
5//! random point. This is done through a layer-by-layer recursive reduction, where each layer uses a
6//! sumcheck protocol.
7//!
8//! The GKR Air Module verifies the [`GkrProof`] struct and
9//! consists of four AIRs:
10//!
11//! 1. **GkrInputAir** - Handles initial setup, coordinates other AIRs, and sends final claims to
12//!    batch constraint module
13//! 2. **GkrLayerAir** - Manages layer-by-layer GKR reduction (verifies
14//!    [`verify_gkr`](openvm_stark_backend::verifier::fractional_sumcheck_gkr::verify_gkr))
15//! 3. **GkrLayerSumcheckAir** - Executes sumcheck protocol for each layer (verifies
16//!    `verify_gkr_sumcheck`)
17//! 4. **GkrXiSamplerAir** - Samples additional xi randomness challenges if required
18//!
19//! ## Architecture
20//!
21//! ```text
22//!                                ┌─────────────────┐
23//!                                │                 │───────────────────► TranscriptBus
24//!                                │ GkrXiSamplerAir │
25//!                                │                 │───────────────────► XiRandomnessBus
26//!                                └─────────────────┘
27//!                                         ▲
28//!                                         ┆
29//!                         GkrXiSamplerBus ┆
30//!                                         ┆
31//!                                         ▼
32//!                                ┌─────────────────┐
33//!                                │                 │───────────────────► TranscriptBus
34//!                                │                 │
35//!  GkrModuleBus ────────────────►│   GkrInputAir   │───────────────────► ExpBitsLenBus
36//!                                │                 │
37//!                                │                 │───────────────────► BatchConstraintModuleBus
38//!                                └─────────────────┘
39//!                                      ┆      ▲
40//!                                      ┆      ┆
41//!                     GkrLayerInputBus ┆      ┆ GkrLayerOutputBus
42//!                                      ┆      ┆
43//!                                      ▼      ┆
44//!                             ┌─────────────────────────┐
45//!                             │                         │──────────────► TranscriptBus
46//!   ┌┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄│       GkrLayerAir       │
47//!   ┆                         │                         │──────────────► XiRandomnessBus
48//!   ┆                         └─────────────────────────┘
49//!   ┆                                  ┆      ▲
50//!   ┆                                  ┆      ┆
51//!   ┆              GkrSumcheckInputBus ┆      ┆ GkrSumcheckOutputBus
52//!   ┆                                  ┆      ┆
53//!   ┆                                  ▼      ┆
54//!   ┆ GkrSumcheckChallengeBus ┌─────────────────────────┐
55//!   ┆┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄│                         │──────────────► TranscriptBus
56//!   ┆                         │   GkrLayerSumcheckAir   │
57//!   └┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄►│                         │──────────────► XiRandomnessBus
58//!                             └─────────────────────────┘
59//! ```
60
61use core::iter::zip;
62use std::sync::Arc;
63
64use itertools::Itertools;
65use openvm_cpu_backend::CpuBackend;
66use openvm_stark_backend::{
67    keygen::types::MultiStarkVerifyingKey,
68    p3_maybe_rayon::prelude::*,
69    poly_common::{interpolate_cubic_at_0123, interpolate_linear_at_01},
70    proof::{GkrProof, Proof},
71    prover::AirProvingContext,
72    AirRef, FiatShamirTranscript, ReadOnlyTranscript, StarkProtocolConfig, TranscriptHistory,
73};
74use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, D_EF, EF, F};
75use p3_field::{Field, PrimeCharacteristicRing};
76use p3_matrix::dense::RowMajorMatrix;
77use strum::EnumCount;
78
79#[cfg(feature = "cuda")]
80use crate::primitives::exp_bits_len::ExpBitsLenTraceGenerator as GpuExpBitsLenTraceGenerator;
81use crate::{
82    gkr::{
83        bus::{GkrLayerInputBus, GkrLayerOutputBus, GkrXiSamplerBus},
84        input::{GkrInputAir, GkrInputRecord, GkrInputTraceGenerator},
85        layer::{GkrLayerAir, GkrLayerRecord, GkrLayerTraceGenerator},
86        sumcheck::{GkrLayerSumcheckAir, GkrSumcheckRecord, GkrSumcheckTraceGenerator},
87        xi_sampler::{GkrXiSamplerAir, GkrXiSamplerRecord, GkrXiSamplerTraceGenerator},
88    },
89    primitives::exp_bits_len::ExpBitsLenCpuTraceGenerator,
90    system::{
91        AirModule, BusIndexManager, BusInventory, GkrPreflight, GlobalCtxCpu, Preflight,
92        TraceGenModule,
93    },
94    tracegen::{ModuleChip, RowMajorChip},
95    utils::{pow_observe_sample, pow_tidx_count},
96};
97
98// Internal bus definitions
99mod bus;
100pub use bus::{
101    GkrSumcheckChallengeBus, GkrSumcheckChallengeMessage, GkrSumcheckInputBus,
102    GkrSumcheckInputMessage, GkrSumcheckOutputBus, GkrSumcheckOutputMessage,
103};
104
105// Sub-modules for different AIRs
106pub mod input;
107pub mod layer;
108pub mod sumcheck;
109pub mod xi_sampler;
110
111pub struct GkrModule {
112    // System Params
113    l_skip: usize,
114    logup_pow_bits: usize,
115    // Global bus inventory
116    bus_inventory: BusInventory,
117    // Module buses
118    xi_sampler_bus: GkrXiSamplerBus,
119    layer_input_bus: GkrLayerInputBus,
120    layer_output_bus: GkrLayerOutputBus,
121    sumcheck_input_bus: GkrSumcheckInputBus,
122    sumcheck_output_bus: GkrSumcheckOutputBus,
123    sumcheck_challenge_bus: GkrSumcheckChallengeBus,
124}
125
126struct GkrBlobCpu {
127    input_records: Vec<GkrInputRecord>,
128    layer_records: Vec<GkrLayerRecord>,
129    sumcheck_records: Vec<GkrSumcheckRecord>,
130    xi_sampler_records: Vec<GkrXiSamplerRecord>,
131    mus_records: Vec<Vec<EF>>,
132    q0_claims: Vec<EF>,
133}
134
135impl GkrModule {
136    pub fn new(
137        mvk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
138        b: &mut BusIndexManager,
139        bus_inventory: BusInventory,
140    ) -> Self {
141        GkrModule {
142            l_skip: mvk.inner.params.l_skip,
143            logup_pow_bits: mvk.inner.params.logup.pow_bits,
144            bus_inventory,
145            layer_input_bus: GkrLayerInputBus::new(b.new_bus_idx()),
146            layer_output_bus: GkrLayerOutputBus::new(b.new_bus_idx()),
147            sumcheck_input_bus: GkrSumcheckInputBus::new(b.new_bus_idx()),
148            sumcheck_output_bus: GkrSumcheckOutputBus::new(b.new_bus_idx()),
149            sumcheck_challenge_bus: GkrSumcheckChallengeBus::new(b.new_bus_idx()),
150            xi_sampler_bus: GkrXiSamplerBus::new(b.new_bus_idx()),
151        }
152    }
153
154    #[tracing::instrument(level = "trace", skip_all)]
155    pub fn run_preflight<TS>(
156        &self,
157        proof: &Proof<BabyBearPoseidon2Config>,
158        preflight: &mut Preflight,
159        ts: &mut TS,
160    ) where
161        TS: FiatShamirTranscript<BabyBearPoseidon2Config> + TranscriptHistory,
162    {
163        let GkrProof {
164            q0_claim,
165            claims_per_layer,
166            sumcheck_polys,
167            logup_pow_witness,
168        } = &proof.gkr_proof;
169
170        let _logup_pow_sample = pow_observe_sample(ts, self.logup_pow_bits, *logup_pow_witness);
171        let _alpha_logup = ts.sample_ext();
172        let _beta_logup = ts.sample_ext();
173
174        let mut xi = vec![(0, EF::ZERO); claims_per_layer.len()];
175        let mut gkr_r = vec![EF::ZERO];
176        let mut numer_claim = EF::ZERO;
177        let mut denom_claim = EF::ONE;
178
179        if !claims_per_layer.is_empty() {
180            debug_assert_eq!(sumcheck_polys.len() + 1, claims_per_layer.len());
181
182            ts.observe_ext(*q0_claim);
183
184            let claims = &claims_per_layer[0];
185
186            ts.observe_ext(claims.p_xi_0);
187            ts.observe_ext(claims.q_xi_0);
188            ts.observe_ext(claims.p_xi_1);
189            ts.observe_ext(claims.q_xi_1);
190
191            let mu = ts.sample_ext();
192            // Reduce layer 0 claims to single evaluation
193            numer_claim = interpolate_linear_at_01(&[claims.p_xi_0, claims.p_xi_1], mu);
194            denom_claim = interpolate_linear_at_01(&[claims.q_xi_0, claims.q_xi_1], mu);
195            gkr_r = vec![mu];
196        }
197
198        for (i, (polys, claims)) in zip(sumcheck_polys, claims_per_layer.iter().skip(1)).enumerate()
199        {
200            let layer_idx = i + 1;
201            let is_final_layer = i == sumcheck_polys.len() - 1;
202
203            let lambda = ts.sample_ext();
204
205            // Compute initial claim for this layer using numer_claim and denom_claim from previous
206            // layer
207            let mut claim = numer_claim + lambda * denom_claim;
208            let mut eq = EF::ONE;
209            let mut gkr_r_prime = Vec::with_capacity(layer_idx);
210
211            for (j, poly) in polys.iter().enumerate() {
212                for eval in poly {
213                    ts.observe_ext(*eval);
214                }
215                let ri = ts.sample_ext();
216
217                // Compute claim_out via cubic interpolation
218                let ev0 = claim - poly[0];
219                let evals = [ev0, poly[0], poly[1], poly[2]];
220                let claim_out = interpolate_cubic_at_0123(&evals, ri);
221
222                // Update eq incrementally: eq *= xi * ri + (1 - xi) * (1 - ri)
223                let xi_j = gkr_r[j];
224                let eq_out = eq * (xi_j * ri + (EF::ONE - xi_j) * (EF::ONE - ri));
225
226                claim = claim_out;
227                eq = eq_out;
228                gkr_r_prime.push(ri);
229
230                if is_final_layer {
231                    xi[j + 1] = (ts.len() - D_EF, ri);
232                }
233            }
234
235            ts.observe_ext(claims.p_xi_0);
236            ts.observe_ext(claims.q_xi_0);
237            ts.observe_ext(claims.p_xi_1);
238            ts.observe_ext(claims.q_xi_1);
239
240            let mu = ts.sample_ext();
241            // Reduce current layer claims to single evaluation for next layer
242            numer_claim = interpolate_linear_at_01(&[claims.p_xi_0, claims.p_xi_1], mu);
243            denom_claim = interpolate_linear_at_01(&[claims.q_xi_0, claims.q_xi_1], mu);
244            gkr_r = std::iter::once(mu).chain(gkr_r_prime).collect();
245
246            if is_final_layer {
247                xi[0] = (ts.len() - D_EF, mu);
248            }
249        }
250
251        for _ in claims_per_layer.len()..preflight.proof_shape.n_max + self.l_skip {
252            xi.push((ts.len(), ts.sample_ext()));
253        }
254
255        preflight.gkr = GkrPreflight {
256            post_tidx: ts.len(),
257            xi,
258        };
259    }
260}
261
262impl AirModule for GkrModule {
263    fn num_airs(&self) -> usize {
264        GkrModuleChipDiscriminants::COUNT
265    }
266
267    fn airs<SC: StarkProtocolConfig<F = F>>(&self) -> Vec<AirRef<SC>> {
268        let gkr_input_air = GkrInputAir {
269            l_skip: self.l_skip,
270            logup_pow_bits: self.logup_pow_bits,
271            gkr_module_bus: self.bus_inventory.gkr_module_bus,
272            bc_module_bus: self.bus_inventory.bc_module_bus,
273            transcript_bus: self.bus_inventory.transcript_bus,
274            exp_bits_len_bus: self.bus_inventory.exp_bits_len_bus,
275            layer_input_bus: self.layer_input_bus,
276            layer_output_bus: self.layer_output_bus,
277            xi_sampler_bus: self.xi_sampler_bus,
278            constraints_folding_input_bus: self.bus_inventory.constraints_folding_input_bus,
279            interactions_folding_input_bus: self.bus_inventory.interactions_folding_input_bus,
280        };
281
282        let gkr_layer_air = GkrLayerAir {
283            xi_randomness_bus: self.bus_inventory.xi_randomness_bus,
284            transcript_bus: self.bus_inventory.transcript_bus,
285            layer_input_bus: self.layer_input_bus,
286            layer_output_bus: self.layer_output_bus,
287            sumcheck_input_bus: self.sumcheck_input_bus,
288            sumcheck_challenge_bus: self.sumcheck_challenge_bus,
289            sumcheck_output_bus: self.sumcheck_output_bus,
290        };
291
292        let gkr_sumcheck_air = GkrLayerSumcheckAir::new(
293            self.bus_inventory.transcript_bus,
294            self.bus_inventory.xi_randomness_bus,
295            self.sumcheck_input_bus,
296            self.sumcheck_output_bus,
297            self.sumcheck_challenge_bus,
298        );
299
300        let gkr_xi_sampler_air = GkrXiSamplerAir {
301            xi_randomness_bus: self.bus_inventory.xi_randomness_bus,
302            transcript_bus: self.bus_inventory.transcript_bus,
303            xi_sampler_bus: self.xi_sampler_bus,
304        };
305
306        vec![
307            Arc::new(gkr_input_air) as AirRef<_>,
308            Arc::new(gkr_layer_air) as AirRef<_>,
309            Arc::new(gkr_sumcheck_air) as AirRef<_>,
310            Arc::new(gkr_xi_sampler_air) as AirRef<_>,
311        ]
312    }
313}
314
315impl GkrModule {
316    #[tracing::instrument(skip_all)]
317    fn generate_blob(
318        &self,
319        _child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
320        proofs: &[&Proof<BabyBearPoseidon2Config>],
321        preflights: &[&Preflight],
322        exp_bits_len_gen: &ExpBitsLenCpuTraceGenerator,
323    ) -> GkrBlobCpu {
324        debug_assert_eq!(proofs.len(), preflights.len());
325
326        // NOTE: we only collect the zipped vec because rayon vs itertools has different treatment
327        // of multiunzip. This could be addressed with a macro similar to parizip!
328        let zipped_records: Vec<_> = proofs
329            .par_iter()
330            .zip(preflights.par_iter())
331            .map(|(proof, preflight)| {
332                let start_idx = preflight.proof_shape.post_tidx;
333                let mut ts = ReadOnlyTranscript::new(&preflight.transcript, start_idx);
334
335                let gkr_proof = &proof.gkr_proof;
336                let GkrProof {
337                    q0_claim,
338                    claims_per_layer,
339                    sumcheck_polys,
340                    logup_pow_witness,
341                } = gkr_proof;
342
343                let logup_pow_sample =
344                    pow_observe_sample(&mut ts, self.logup_pow_bits, *logup_pow_witness);
345                if self.logup_pow_bits > 0 {
346                    exp_bits_len_gen.add_request(
347                        F::GENERATOR,
348                        logup_pow_sample,
349                        self.logup_pow_bits,
350                    );
351                }
352
353                let alpha_logup =
354                    FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
355                let _beta_logup =
356                    FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
357
358                let xi = &preflight.gkr.xi;
359
360                let input_layer_claim = claims_per_layer
361                    .last()
362                    .and_then(|last_layer| {
363                        xi.first().map(|(_, rho)| {
364                            let p_claim =
365                                last_layer.p_xi_0 + *rho * (last_layer.p_xi_1 - last_layer.p_xi_0);
366                            let q_claim =
367                                last_layer.q_xi_0 + *rho * (last_layer.q_xi_1 - last_layer.q_xi_0);
368                            [p_claim, q_claim]
369                        })
370                    })
371                    .unwrap_or([EF::ZERO, alpha_logup]);
372
373                let input_record = GkrInputRecord {
374                    tidx: preflight.proof_shape.post_tidx,
375                    n_logup: preflight.proof_shape.n_logup,
376                    n_max: preflight.proof_shape.n_max,
377                    logup_pow_witness: *logup_pow_witness,
378                    logup_pow_sample,
379                    alpha_logup,
380                    input_layer_claim,
381                };
382
383                let num_layers = claims_per_layer.len();
384                let sumcheck_layer_count = sumcheck_polys.len();
385                let total_sumcheck_rounds: usize = sumcheck_polys.iter().map(Vec::len).sum();
386
387                let logup_pow_offset = pow_tidx_count(self.logup_pow_bits);
388                let tidx_first_gkr_layer =
389                    preflight.proof_shape.post_tidx + logup_pow_offset + 2 * D_EF + D_EF;
390                let mut layer_record = GkrLayerRecord {
391                    tidx: tidx_first_gkr_layer,
392                    layer_claims: Vec::with_capacity(num_layers),
393                    lambdas: Vec::with_capacity(sumcheck_layer_count),
394                    eq_at_r_primes: Vec::with_capacity(sumcheck_layer_count),
395                };
396                let mut mus = Vec::with_capacity(num_layers.max(1));
397
398                let tidx_first_sumcheck_round = tidx_first_gkr_layer + 5 * D_EF + D_EF;
399                let mut sumcheck_record = GkrSumcheckRecord {
400                    tidx: tidx_first_sumcheck_round,
401                    ris: Vec::with_capacity(total_sumcheck_rounds),
402                    evals: Vec::with_capacity(total_sumcheck_rounds),
403                    claims: Vec::with_capacity(sumcheck_layer_count),
404                };
405
406                let mut gkr_r: Vec<EF> = Vec::new();
407                let mut numer_claim = EF::ZERO;
408                let mut denom_claim = EF::ONE;
409
410                if let Some(root_claims) = claims_per_layer.first() {
411                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
412                        &mut ts, *q0_claim,
413                    );
414                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
415                        &mut ts,
416                        root_claims.p_xi_0,
417                    );
418                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
419                        &mut ts,
420                        root_claims.q_xi_0,
421                    );
422                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
423                        &mut ts,
424                        root_claims.p_xi_1,
425                    );
426                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
427                        &mut ts,
428                        root_claims.q_xi_1,
429                    );
430
431                    let mu = FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
432                    numer_claim =
433                        interpolate_linear_at_01(&[root_claims.p_xi_0, root_claims.p_xi_1], mu);
434                    denom_claim =
435                        interpolate_linear_at_01(&[root_claims.q_xi_0, root_claims.q_xi_1], mu);
436
437                    gkr_r.push(mu);
438
439                    layer_record.layer_claims.push([
440                        root_claims.p_xi_0,
441                        root_claims.q_xi_0,
442                        root_claims.p_xi_1,
443                        root_claims.q_xi_1,
444                    ]);
445                    mus.push(mu);
446                }
447
448                for (polys, claims) in sumcheck_polys.iter().zip(claims_per_layer.iter().skip(1)) {
449                    let lambda =
450                        FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
451                    layer_record.lambdas.push(lambda);
452
453                    let mut claim = numer_claim + lambda * denom_claim;
454                    let mut eq_at_r_prime = EF::ONE;
455                    let mut round_r = Vec::with_capacity(polys.len());
456
457                    sumcheck_record.claims.push(claim);
458
459                    for (round_idx, poly) in polys.iter().enumerate() {
460                        for eval in poly {
461                            FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
462                                &mut ts, *eval,
463                            );
464                        }
465
466                        let ri =
467                            FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
468                        let prev_challenge = gkr_r[round_idx];
469
470                        let ev0 = claim - poly[0];
471                        let evals = [ev0, poly[0], poly[1], poly[2]];
472                        claim = interpolate_cubic_at_0123(&evals, ri);
473
474                        let eq_factor =
475                            prev_challenge * ri + (EF::ONE - prev_challenge) * (EF::ONE - ri);
476                        eq_at_r_prime *= eq_factor;
477
478                        sumcheck_record.ris.push(ri);
479                        sumcheck_record.evals.push(*poly);
480                        round_r.push(ri);
481                    }
482
483                    layer_record.eq_at_r_primes.push(eq_at_r_prime);
484
485                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
486                        &mut ts,
487                        claims.p_xi_0,
488                    );
489                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
490                        &mut ts,
491                        claims.q_xi_0,
492                    );
493                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
494                        &mut ts,
495                        claims.p_xi_1,
496                    );
497                    FiatShamirTranscript::<BabyBearPoseidon2Config>::observe_ext(
498                        &mut ts,
499                        claims.q_xi_1,
500                    );
501
502                    let mu = FiatShamirTranscript::<BabyBearPoseidon2Config>::sample_ext(&mut ts);
503                    numer_claim = interpolate_linear_at_01(&[claims.p_xi_0, claims.p_xi_1], mu);
504                    denom_claim = interpolate_linear_at_01(&[claims.q_xi_0, claims.q_xi_1], mu);
505
506                    gkr_r.clear();
507                    gkr_r.push(mu);
508                    gkr_r.extend(round_r);
509
510                    layer_record.layer_claims.push([
511                        claims.p_xi_0,
512                        claims.q_xi_0,
513                        claims.p_xi_1,
514                        claims.q_xi_1,
515                    ]);
516                    mus.push(mu);
517                }
518
519                let xi_sampler_record = if num_layers < xi.len() {
520                    let challenges: Vec<EF> =
521                        xi.iter().skip(num_layers).map(|(_, val)| *val).collect();
522                    let tidx = xi[num_layers].0;
523                    GkrXiSamplerRecord {
524                        tidx,
525                        idx: num_layers,
526                        xis: challenges,
527                    }
528                } else {
529                    GkrXiSamplerRecord::default()
530                };
531
532                (
533                    input_record,
534                    layer_record,
535                    sumcheck_record,
536                    xi_sampler_record,
537                    mus,
538                    *q0_claim,
539                )
540            })
541            .collect();
542        let (
543            input_records,
544            layer_records,
545            sumcheck_records,
546            xi_sampler_records,
547            mus_records,
548            q0_claims,
549        ): (Vec<_>, Vec<_>, Vec<_>, Vec<_>, Vec<_>, Vec<_>) =
550            zipped_records.into_iter().multiunzip();
551
552        GkrBlobCpu {
553            input_records,
554            layer_records,
555            sumcheck_records,
556            xi_sampler_records,
557            mus_records,
558            q0_claims,
559        }
560    }
561}
562
563impl<SC: StarkProtocolConfig<F = F>> TraceGenModule<GlobalCtxCpu, CpuBackend<SC>> for GkrModule {
564    type ModuleSpecificCtx<'a> = ExpBitsLenCpuTraceGenerator;
565
566    #[tracing::instrument(skip_all)]
567    fn generate_proving_ctxs(
568        &self,
569        child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
570        proofs: &[Proof<BabyBearPoseidon2Config>],
571        preflights: &[Preflight],
572        exp_bits_len_gen: &ExpBitsLenCpuTraceGenerator,
573        required_heights: Option<&[usize]>,
574    ) -> Option<Vec<AirProvingContext<CpuBackend<SC>>>> {
575        let proof_refs = proofs.iter().collect_vec();
576        let preflight_refs = preflights.iter().collect_vec();
577        let blob = self.generate_blob(child_vk, &proof_refs, &preflight_refs, exp_bits_len_gen);
578
579        let chips = [
580            GkrModuleChip::Input,
581            GkrModuleChip::Layer,
582            GkrModuleChip::LayerSumcheck,
583            GkrModuleChip::XiSampler,
584        ];
585
586        let span = tracing::Span::current();
587        chips
588            .par_iter()
589            .map(|chip| {
590                let _guard = span.enter();
591                chip.generate_proving_ctx(
592                    &blob,
593                    required_heights.map(|heights| heights[chip.index()]),
594                )
595            })
596            .collect::<Vec<_>>()
597            .into_iter()
598            .collect()
599    }
600}
601
602// To reduce the number of structs and trait implementations, we collect them into a single enum
603// with enum dispatch.
604#[derive(strum_macros::Display, strum::EnumDiscriminants)]
605#[strum_discriminants(derive(strum_macros::EnumCount))]
606#[strum_discriminants(repr(usize))]
607enum GkrModuleChip {
608    Input,
609    Layer,
610    LayerSumcheck,
611    XiSampler,
612}
613
614impl GkrModuleChip {
615    fn index(&self) -> usize {
616        GkrModuleChipDiscriminants::from(self) as usize
617    }
618}
619
620impl RowMajorChip<F> for GkrModuleChip {
621    type Ctx<'a> = GkrBlobCpu;
622
623    #[tracing::instrument(
624        name = "wrapper.generate_trace",
625        level = "trace",
626        skip_all,
627        fields(air = %self)
628    )]
629    fn generate_trace(
630        &self,
631        blob: &Self::Ctx<'_>,
632        required_height: Option<usize>,
633    ) -> Option<RowMajorMatrix<F>> {
634        use GkrModuleChip::*;
635        match self {
636            Input => GkrInputTraceGenerator
637                .generate_trace(&(&blob.input_records, &blob.q0_claims), required_height),
638            Layer => GkrLayerTraceGenerator.generate_trace(
639                &(&blob.layer_records, &blob.mus_records, &blob.q0_claims),
640                required_height,
641            ),
642            LayerSumcheck => GkrSumcheckTraceGenerator.generate_trace(
643                &(&blob.sumcheck_records, &blob.mus_records),
644                required_height,
645            ),
646            XiSampler => GkrXiSamplerTraceGenerator
647                .generate_trace(&blob.xi_sampler_records.as_slice(), required_height),
648        }
649    }
650}
651
652#[cfg(feature = "cuda")]
653mod cuda_tracegen {
654    use itertools::Itertools;
655    use openvm_cuda_backend::{data_transporter::transport_matrix_h2d_row, GpuBackend};
656    use openvm_cuda_common::stream::GpuDeviceCtx;
657    use openvm_stark_backend::{p3_maybe_rayon::prelude::*, prover::AirProvingContext};
658
659    use super::*;
660    use crate::cuda::{
661        preflight::PreflightGpu, proof::ProofGpu, vk::VerifyingKeyGpu, GlobalCtxGpu,
662    };
663
664    impl TraceGenModule<GlobalCtxGpu, GpuBackend> for GkrModule {
665        type ModuleSpecificCtx<'a> = (&'a GpuExpBitsLenTraceGenerator, &'a GpuDeviceCtx);
666
667        #[tracing::instrument(skip_all)]
668        fn generate_proving_ctxs(
669            &self,
670            child_vk: &VerifyingKeyGpu,
671            proofs: &[ProofGpu],
672            preflights: &[PreflightGpu],
673            module_ctx: &Self::ModuleSpecificCtx<'_>,
674            required_heights: Option<&[usize]>,
675        ) -> Option<Vec<AirProvingContext<GpuBackend>>> {
676            let exp_bits_len_gen = module_ctx.0;
677            let device_ctx = module_ctx.1;
678            let proofs_cpu = proofs.iter().map(|proof| &proof.cpu).collect_vec();
679            let preflights_cpu = preflights
680                .iter()
681                .map(|preflight| &preflight.cpu)
682                .collect_vec();
683            let blob = self.generate_blob(
684                &child_vk.cpu,
685                &proofs_cpu,
686                &preflights_cpu,
687                exp_bits_len_gen,
688            );
689            let chips = [
690                GkrModuleChip::Input,
691                GkrModuleChip::Layer,
692                GkrModuleChip::LayerSumcheck,
693                GkrModuleChip::XiSampler,
694            ];
695
696            // Phase 1: CPU trace generation in parallel
697            let span = tracing::Span::current();
698            let cpu_traces: Vec<_> = chips
699                .par_iter()
700                .map(|chip| {
701                    let _guard = span.enter();
702                    chip.generate_trace(
703                        &blob,
704                        required_heights.map(|heights| heights[chip.index()]),
705                    )
706                })
707                .collect();
708
709            // Phase 2: H2D transfer serially on main thread
710            cpu_traces
711                .into_iter()
712                .map(|trace| {
713                    trace.map(|m| {
714                        AirProvingContext::simple_no_pis(
715                            transport_matrix_h2d_row(&m, device_ctx).unwrap(),
716                        )
717                    })
718                })
719                .collect()
720        }
721    }
722}