1use 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
98mod bus;
100pub use bus::{
101 GkrSumcheckChallengeBus, GkrSumcheckChallengeMessage, GkrSumcheckInputBus,
102 GkrSumcheckInputMessage, GkrSumcheckOutputBus, GkrSumcheckOutputMessage,
103};
104
105pub mod input;
107pub mod layer;
108pub mod sumcheck;
109pub mod xi_sampler;
110
111pub struct GkrModule {
112 l_skip: usize,
114 logup_pow_bits: usize,
115 bus_inventory: BusInventory,
117 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 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 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 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 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 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 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#[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 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 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}