openvm_recursion_circuit/stacking/
mod.rs

1use std::sync::Arc;
2
3use itertools::{izip, Itertools};
4use openvm_cpu_backend::CpuBackend;
5use openvm_stark_backend::{
6    keygen::types::MultiStarkVerifyingKey,
7    p3_maybe_rayon::prelude::*,
8    proof::{Proof, StackingProof},
9    prover::AirProvingContext,
10    AirRef, FiatShamirTranscript, StarkProtocolConfig, TranscriptHistory,
11};
12use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, F};
13use p3_field::PrimeCharacteristicRing;
14use p3_matrix::dense::RowMajorMatrix;
15use strum::{EnumCount, EnumDiscriminants};
16
17use crate::{
18    stacking::{
19        bus::*,
20        claims::{StackingClaimsAir, StackingClaimsTraceGenerator},
21        eq_base::{EqBaseAir, EqBaseTraceGenerator},
22        eq_bits::{EqBitsAir, EqBitsTraceGenerator},
23        opening::{OpeningClaimsAir, OpeningClaimsTraceGenerator},
24        sumcheck::{SumcheckRoundsAir, SumcheckRoundsTraceGenerator},
25        univariate::{UnivariateRoundAir, UnivariateRoundTraceGenerator},
26    },
27    system::{
28        AirModule, BusIndexManager, BusInventory, GlobalCtxCpu, Preflight, StackingPreflight,
29        TraceGenModule,
30    },
31    tracegen::{ModuleChip, RowMajorChip, StandardTracegenCtx},
32    utils::pow_observe_sample,
33};
34
35mod bus;
36pub mod claims;
37pub mod eq_base;
38pub mod eq_bits;
39pub mod opening;
40pub mod sumcheck;
41pub mod univariate;
42mod utils;
43
44#[cfg(feature = "cuda")]
45mod cuda_abi;
46
47pub struct StackingModule {
48    bus_inventory: BusInventory,
49
50    // Internal buses
51    stacking_tidx_bus: StackingModuleTidxBus,
52    claim_coefficients_bus: ClaimCoefficientsBus,
53    sumcheck_claims_bus: SumcheckClaimsBus,
54    eq_rand_values_bus: EqRandValuesLookupBus,
55    eq_base_bus: EqBaseBus,
56    eq_bits_internal_bus: EqBitsInternalBus,
57    eq_kernel_lookup_bus: EqKernelLookupBus,
58    eq_bits_lookup_bus: EqBitsLookupBus,
59
60    l_skip: usize,
61    n_stack: usize,
62    w_stack: usize,
63    stacking_index_mult: usize,
64    /// Number of PoW bits for μ batching challenge.
65    mu_pow_bits: usize,
66}
67
68impl StackingModule {
69    pub fn new(
70        child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
71        b: &mut BusIndexManager,
72        bus_inventory: BusInventory,
73    ) -> Self {
74        Self {
75            bus_inventory,
76            stacking_tidx_bus: StackingModuleTidxBus::new(b.new_bus_idx()),
77            claim_coefficients_bus: ClaimCoefficientsBus::new(b.new_bus_idx()),
78            sumcheck_claims_bus: SumcheckClaimsBus::new(b.new_bus_idx()),
79            eq_rand_values_bus: EqRandValuesLookupBus::new(b.new_bus_idx()),
80            eq_base_bus: EqBaseBus::new(b.new_bus_idx()),
81            eq_bits_internal_bus: EqBitsInternalBus::new(b.new_bus_idx()),
82            eq_kernel_lookup_bus: EqKernelLookupBus::new(b.new_bus_idx()),
83            eq_bits_lookup_bus: EqBitsLookupBus::new(b.new_bus_idx()),
84            l_skip: child_vk.inner.params.l_skip,
85            n_stack: child_vk.inner.params.n_stack,
86            w_stack: child_vk.inner.params.w_stack,
87            stacking_index_mult: child_vk
88                .inner
89                .params
90                .whir
91                .rounds
92                .first()
93                .map(|round| round.num_queries)
94                .unwrap_or(0)
95                << child_vk.inner.params.k_whir(),
96            mu_pow_bits: child_vk.inner.params.whir.mu_pow_bits,
97        }
98    }
99
100    #[tracing::instrument(level = "trace", skip_all)]
101    pub fn run_preflight<TS>(
102        &self,
103        proof: &Proof<BabyBearPoseidon2Config>,
104        preflight: &mut Preflight,
105        ts: &mut TS,
106    ) where
107        TS: FiatShamirTranscript<BabyBearPoseidon2Config> + TranscriptHistory,
108    {
109        let mut sumcheck_rnd = vec![];
110        let mut intermediate_tidx = [0; 3];
111
112        let StackingProof {
113            univariate_round_coeffs,
114            sumcheck_round_polys,
115            stacking_openings,
116        } = &proof.stacking_proof;
117
118        let lambda = ts.sample_ext();
119        intermediate_tidx[0] = ts.len();
120
121        for coef in univariate_round_coeffs {
122            ts.observe_ext(*coef);
123        }
124        let u0 = ts.sample_ext();
125        let univariate_poly_rand_eval = izip!(univariate_round_coeffs, u0.powers())
126            .map(|(&coef, pow)| coef * pow)
127            .sum();
128        sumcheck_rnd.push(u0);
129        intermediate_tidx[1] = ts.len();
130
131        for poly in sumcheck_round_polys {
132            for eval in poly {
133                ts.observe_ext(*eval);
134            }
135            let ui = ts.sample_ext();
136            sumcheck_rnd.push(ui);
137        }
138        intermediate_tidx[2] = ts.len();
139
140        for matrix_openings in stacking_openings {
141            for col_opening in matrix_openings {
142                ts.observe_ext(*col_opening);
143            }
144        }
145
146        // μ PoW: observe witness and sample before sampling μ
147        let mu_pow_witness = proof.whir_proof.mu_pow_witness;
148        let mu_pow_sample = pow_observe_sample(ts, self.mu_pow_bits, mu_pow_witness);
149
150        let stacking_batching_challenge = ts.sample_ext();
151
152        preflight.stacking = StackingPreflight {
153            intermediate_tidx,
154            post_tidx: ts.len(),
155            univariate_poly_rand_eval,
156            stacking_batching_challenge,
157            mu_pow_witness,
158            mu_pow_sample,
159            lambda,
160            sumcheck_rnd,
161        };
162    }
163}
164
165impl AirModule for StackingModule {
166    fn num_airs(&self) -> usize {
167        StackingModuleChipDiscriminants::COUNT
168    }
169
170    fn airs<SC: StarkProtocolConfig<F = F>>(&self) -> Vec<AirRef<SC>> {
171        let opening_air = OpeningClaimsAir {
172            lifted_heights_bus: self.bus_inventory.lifted_heights_bus,
173            stacking_module_bus: self.bus_inventory.stacking_module_bus,
174            column_claims_bus: self.bus_inventory.column_claims_bus,
175            transcript_bus: self.bus_inventory.transcript_bus,
176            air_shape_bus: self.bus_inventory.air_shape_bus,
177            stacking_tidx_bus: self.stacking_tidx_bus,
178            claim_coefficients_bus: self.claim_coefficients_bus,
179            sumcheck_claims_bus: self.sumcheck_claims_bus,
180            eq_kernel_lookup_bus: self.eq_kernel_lookup_bus,
181            eq_bits_lookup_bus: self.eq_bits_lookup_bus,
182            l_skip: self.l_skip,
183            n_stack: self.n_stack,
184        };
185        let univariate_round_air = UnivariateRoundAir {
186            transcript_bus: self.bus_inventory.transcript_bus,
187            stacking_tidx_bus: self.stacking_tidx_bus,
188            sumcheck_claims_bus: self.sumcheck_claims_bus,
189            eq_rand_values_bus: self.eq_rand_values_bus,
190            eq_kernel_lookup_bus: self.eq_kernel_lookup_bus,
191            l_skip: self.l_skip,
192        };
193        let sumcheck_rounds_air = SumcheckRoundsAir {
194            constraint_randomness_bus: self.bus_inventory.constraint_randomness_bus,
195            whir_opening_point_bus: self.bus_inventory.whir_opening_point_bus,
196            transcript_bus: self.bus_inventory.transcript_bus,
197            stacking_tidx_bus: self.stacking_tidx_bus,
198            sumcheck_claims_bus: self.sumcheck_claims_bus,
199            eq_base_bus: self.eq_base_bus,
200            eq_rand_values_bus: self.eq_rand_values_bus,
201            eq_kernel_lookup_bus: self.eq_kernel_lookup_bus,
202            l_skip: self.l_skip,
203        };
204        let stacking_claims_air = StackingClaimsAir {
205            stacking_indices_bus: self.bus_inventory.stacking_indices_bus,
206            whir_module_bus: self.bus_inventory.whir_module_bus,
207            whir_mu_bus: self.bus_inventory.whir_mu_bus,
208            transcript_bus: self.bus_inventory.transcript_bus,
209            exp_bits_len_bus: self.bus_inventory.exp_bits_len_bus,
210            stacking_tidx_bus: self.stacking_tidx_bus,
211            claim_coefficients_bus: self.claim_coefficients_bus,
212            sumcheck_claims_bus: self.sumcheck_claims_bus,
213            stacking_index_mult: self.stacking_index_mult,
214            w_stack: self.w_stack,
215            mu_pow_bits: self.mu_pow_bits,
216        };
217        let eq_base_air = EqBaseAir {
218            constraint_randomness_bus: self.bus_inventory.constraint_randomness_bus,
219            whir_opening_point_bus: self.bus_inventory.whir_opening_point_bus,
220            eq_base_bus: self.eq_base_bus,
221            eq_rand_values_bus: self.eq_rand_values_bus,
222            eq_kernel_lookup_bus: self.eq_kernel_lookup_bus,
223            eq_neg_base_rand_bus: self.bus_inventory.eq_neg_base_rand_bus,
224            eq_neg_result_bus: self.bus_inventory.eq_neg_result_bus,
225            l_skip: self.l_skip,
226        };
227        let eq_bits_air = EqBitsAir {
228            eq_bits_internal_bus: self.eq_bits_internal_bus,
229            eq_bits_lookup_bus: self.eq_bits_lookup_bus,
230            eq_rand_values_bus: self.eq_rand_values_bus,
231            n_stack: self.n_stack,
232            l_skip: self.l_skip,
233        };
234        vec![
235            Arc::new(opening_air),
236            Arc::new(univariate_round_air),
237            Arc::new(sumcheck_rounds_air),
238            Arc::new(stacking_claims_air),
239            Arc::new(eq_base_air),
240            Arc::new(eq_bits_air),
241        ]
242    }
243}
244
245#[derive(Clone, Copy, strum_macros::Display, EnumDiscriminants)]
246#[strum_discriminants(derive(strum_macros::EnumCount))]
247#[strum_discriminants(repr(usize))]
248enum StackingModuleChip {
249    OpeningClaims,
250    UnivariateRound,
251    SumcheckRounds,
252    StackingClaims,
253    EqBase,
254    EqBits,
255}
256
257impl StackingModuleChip {
258    fn index(&self) -> usize {
259        StackingModuleChipDiscriminants::from(self) as usize
260    }
261}
262
263impl RowMajorChip<F> for StackingModuleChip {
264    type Ctx<'a> = StandardTracegenCtx<'a>;
265
266    #[tracing::instrument(
267        name = "wrapper.generate_trace",
268        level = "trace",
269        skip_all,
270        fields(air = %self)
271    )]
272    fn generate_trace(
273        &self,
274        ctx: &Self::Ctx<'_>,
275        required_height: Option<usize>,
276    ) -> Option<RowMajorMatrix<F>> {
277        match self {
278            StackingModuleChip::OpeningClaims => {
279                OpeningClaimsTraceGenerator.generate_trace(ctx, required_height)
280            }
281            StackingModuleChip::UnivariateRound => {
282                UnivariateRoundTraceGenerator.generate_trace(ctx, required_height)
283            }
284            StackingModuleChip::SumcheckRounds => {
285                SumcheckRoundsTraceGenerator.generate_trace(ctx, required_height)
286            }
287            StackingModuleChip::StackingClaims => {
288                StackingClaimsTraceGenerator.generate_trace(ctx, required_height)
289            }
290            StackingModuleChip::EqBase => EqBaseTraceGenerator.generate_trace(ctx, required_height),
291            StackingModuleChip::EqBits => EqBitsTraceGenerator.generate_trace(ctx, required_height),
292        }
293    }
294}
295
296impl<SC: StarkProtocolConfig<F = F>> TraceGenModule<GlobalCtxCpu, CpuBackend<SC>>
297    for StackingModule
298{
299    type ModuleSpecificCtx<'a> = ();
300
301    #[tracing::instrument(skip_all)]
302    fn generate_proving_ctxs(
303        &self,
304        child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
305        proofs: &[Proof<BabyBearPoseidon2Config>],
306        preflights: &[Preflight],
307        _module_ctx: &(),
308        required_heights: Option<&[usize]>,
309    ) -> Option<Vec<AirProvingContext<CpuBackend<SC>>>> {
310        let ctx = StandardTracegenCtx {
311            vk: child_vk,
312            proofs: &proofs.iter().collect_vec(),
313            preflights: &preflights.iter().collect_vec(),
314        };
315        let chips = [
316            StackingModuleChip::OpeningClaims,
317            StackingModuleChip::UnivariateRound,
318            StackingModuleChip::SumcheckRounds,
319            StackingModuleChip::StackingClaims,
320            StackingModuleChip::EqBase,
321            StackingModuleChip::EqBits,
322        ];
323        let span = tracing::Span::current();
324        chips
325            .par_iter()
326            .map(|chip| {
327                let _guard = span.enter();
328                chip.generate_proving_ctx(
329                    &ctx,
330                    required_heights.map(|heights| heights[chip.index()]),
331                )
332            })
333            .collect::<Vec<_>>()
334            .into_iter()
335            .collect()
336    }
337}
338
339#[cfg(feature = "cuda")]
340mod cuda_tracegen {
341    use itertools::Itertools;
342    use openvm_cuda_backend::{
343        data_transporter::transport_matrix_h2d_row, prelude::EF, GpuBackend,
344    };
345    use openvm_cuda_common::{copy::MemCopyH2D, d_buffer::DeviceBuffer, stream::GpuDeviceCtx};
346    use openvm_stark_backend::p3_maybe_rayon::prelude::*;
347
348    use super::*;
349    use crate::{
350        cuda::{preflight::PreflightGpu, proof::ProofGpu, vk::VerifyingKeyGpu, GlobalCtxGpu},
351        stacking::{
352            claims::cuda::StackingClaimsTraceGeneratorGpu,
353            cuda_abi::{
354                compute_coefficients, compute_coefficients_temp_bytes, stacked_slice_data,
355                PolyPrecomputation, StackedSliceData, StackedTraceData,
356            },
357            opening::cuda::OpeningClaimsTraceGeneratorGpu,
358        },
359        tracegen::{cuda::StandardTracegenGpuCtx, RowMajorChip, StandardTracegenCtx},
360    };
361
362    impl ModuleChip<GpuBackend> for StackingModuleChip {
363        type Ctx<'a> = (StandardTracegenGpuCtx<'a>, &'a StackingBlob);
364
365        fn generate_proving_ctx(
366            &self,
367            ctx: &Self::Ctx<'_>,
368            required_height: Option<usize>,
369        ) -> Option<AirProvingContext<GpuBackend>> {
370            match self {
371                StackingModuleChip::OpeningClaims => {
372                    OpeningClaimsTraceGeneratorGpu.generate_proving_ctx(ctx, required_height)
373                }
374                StackingModuleChip::StackingClaims => {
375                    StackingClaimsTraceGeneratorGpu.generate_proving_ctx(ctx, required_height)
376                }
377                _ => {
378                    let proofs_cpu = ctx.0.proofs.iter().map(|p| &p.cpu).collect_vec();
379                    let preflights_cpu = ctx.0.preflights.iter().map(|p| &p.cpu).collect_vec();
380                    let cpu_ctx = StandardTracegenCtx {
381                        vk: &ctx.0.vk.cpu,
382                        proofs: &proofs_cpu,
383                        preflights: &preflights_cpu,
384                    };
385                    let trace = RowMajorChip::generate_trace(self, &cpu_ctx, required_height);
386                    trace.map(|m| {
387                        AirProvingContext::simple_no_pis(
388                            transport_matrix_h2d_row(&m, ctx.0.device_ctx).unwrap(),
389                        )
390                    })
391                }
392            }
393        }
394    }
395
396    pub(crate) struct StackingBlob {
397        pub(crate) slice_data: Vec<DeviceBuffer<StackedSliceData>>,
398        pub(crate) coeffs: Vec<DeviceBuffer<EF>>,
399        pub(crate) precomps: Vec<DeviceBuffer<PolyPrecomputation>>,
400    }
401
402    impl StackingBlob {
403        pub fn new(
404            child_vk: &VerifyingKeyGpu,
405            proofs: &[ProofGpu],
406            preflights: &[PreflightGpu],
407            device_ctx: &GpuDeviceCtx,
408        ) -> Self {
409            let l_skip = child_vk.system_params.l_skip as u32;
410            let n_stack = child_vk.system_params.n_stack as u32;
411            let log_stacked_height = l_skip + n_stack;
412            let stacked_height = 1u32 << log_stacked_height;
413
414            let mut slice_data = Vec::with_capacity(preflights.len());
415            let mut coeffs = Vec::with_capacity(preflights.len());
416            let mut precomps = Vec::with_capacity(preflights.len());
417
418            for (proof, preflight) in izip!(proofs, preflights) {
419                let sorted_trace_data = &preflight.cpu.proof_shape.sorted_trace_vdata;
420                let num_airs = sorted_trace_data.len();
421                let mut num_commits = 1;
422
423                let mut stacked_trace_data = Vec::with_capacity(num_airs);
424                let mut other_trace_data = vec![];
425
426                let mut current_col_idx = 0;
427                let mut current_row_idx = 0;
428
429                for (air_idx, vdata) in sorted_trace_data {
430                    let stark_vk = &child_vk.cpu.inner.per_air[*air_idx];
431                    let trace_widths = &stark_vk.params.width;
432                    let need_rot = stark_vk.params.need_rot as u32;
433                    let log_height = vdata.log_height as u32;
434                    // IMPORTANT: This must match CPU `get_stacked_slice_data`, which stacks ONLY
435                    // the common-main columns here (cached mains are handled as separate commits).
436                    let common_width = trace_widths.common_main as u32;
437
438                    stacked_trace_data.push(StackedTraceData {
439                        commit_idx: 0,
440                        start_col_idx: current_col_idx,
441                        start_row_idx: current_row_idx,
442                        log_height,
443                        width: common_width,
444                        need_rot,
445                    });
446
447                    let lifted_height = log_height.max(l_skip);
448                    let unbounded_row_idx = current_row_idx + (common_width << lifted_height);
449                    current_col_idx += unbounded_row_idx >> log_stacked_height;
450                    current_row_idx = unbounded_row_idx & (stacked_height - 1);
451
452                    other_trace_data.extend(
453                        trace_widths
454                            .preprocessed
455                            .iter()
456                            .chain(trace_widths.cached_mains.iter())
457                            .map(|&part_width| {
458                                let ret = StackedTraceData {
459                                    commit_idx: num_commits,
460                                    start_col_idx: 0,
461                                    start_row_idx: 0,
462                                    log_height,
463                                    width: part_width as u32,
464                                    need_rot,
465                                };
466                                num_commits += 1;
467                                ret
468                            }),
469                    );
470                }
471
472                stacked_trace_data.extend(other_trace_data.iter());
473
474                let mut num_slices = 0u32;
475                let slice_offsets = stacked_trace_data
476                    .iter()
477                    .map(|trace_data| {
478                        let ret = num_slices;
479                        num_slices += trace_data.width;
480                        ret
481                    })
482                    .collect_vec();
483
484                let d_slice_offsets = slice_offsets.to_device_on(device_ctx).unwrap();
485                let d_stacked_trace_data = stacked_trace_data.to_device_on(device_ctx).unwrap();
486
487                let d_slice_data = DeviceBuffer::<StackedSliceData>::with_capacity_on(
488                    num_slices as usize,
489                    device_ctx,
490                );
491
492                unsafe {
493                    stacked_slice_data(
494                        &d_slice_data,
495                        &d_slice_offsets,
496                        &d_stacked_trace_data,
497                        num_airs as u32,
498                        num_commits,
499                        num_slices,
500                        n_stack,
501                        l_skip,
502                        device_ctx.stream.as_raw(),
503                    )
504                    .unwrap();
505                }
506
507                let d_lambda_pows = preflight
508                    .cpu
509                    .stacking
510                    .lambda
511                    .powers()
512                    .take((num_slices << 1) as usize)
513                    .collect_vec()
514                    .to_device_on(device_ctx)
515                    .unwrap();
516
517                let d_coeff_terms =
518                    DeviceBuffer::<EF>::with_capacity_on(num_slices as usize, device_ctx);
519                let d_coeff_term_keys =
520                    DeviceBuffer::<u64>::with_capacity_on(num_slices as usize, device_ctx);
521                let d_precomps = DeviceBuffer::<PolyPrecomputation>::with_capacity_on(
522                    num_slices as usize,
523                    device_ctx,
524                );
525                let d_num_coeffs = DeviceBuffer::<usize>::with_capacity_on(1, device_ctx);
526
527                let num_claims = proof
528                    .cpu
529                    .stacking_proof
530                    .stacking_openings
531                    .iter()
532                    .fold(0, |acc, v| acc + v.len());
533                let d_coeffs = DeviceBuffer::<EF>::with_capacity_on(num_claims, device_ctx);
534                let d_coeff_keys = DeviceBuffer::<u64>::with_capacity_on(num_claims, device_ctx);
535
536                unsafe {
537                    let temp_bytes = compute_coefficients_temp_bytes(
538                        &d_coeff_terms,
539                        &d_coeff_term_keys,
540                        &d_coeffs,
541                        &d_coeff_keys,
542                        num_slices,
543                        &d_num_coeffs,
544                        device_ctx.stream.as_raw(),
545                    )
546                    .unwrap();
547                    let d_temp_buffer =
548                        DeviceBuffer::<u8>::with_capacity_on(temp_bytes, device_ctx);
549                    compute_coefficients(
550                        &d_coeff_terms,
551                        &d_coeff_term_keys,
552                        &d_coeffs,
553                        &d_coeff_keys,
554                        &d_precomps,
555                        &d_slice_data,
556                        &preflight.stacking.sumcheck_rnd,
557                        &preflight.batch_constraint.sumcheck_rnd,
558                        &d_lambda_pows,
559                        num_commits,
560                        num_slices,
561                        n_stack,
562                        l_skip,
563                        &d_temp_buffer,
564                        temp_bytes,
565                        &d_num_coeffs,
566                        device_ctx.stream.as_raw(),
567                    )
568                    .unwrap();
569                }
570
571                slice_data.push(d_slice_data);
572                coeffs.push(d_coeffs);
573                precomps.push(d_precomps);
574            }
575
576            Self {
577                slice_data,
578                coeffs,
579                precomps,
580            }
581        }
582    }
583
584    impl TraceGenModule<GlobalCtxGpu, GpuBackend> for StackingModule {
585        type ModuleSpecificCtx<'a> = openvm_cuda_common::stream::GpuDeviceCtx;
586
587        #[tracing::instrument(skip_all)]
588        fn generate_proving_ctxs(
589            &self,
590            child_vk: &VerifyingKeyGpu,
591            proofs: &[ProofGpu],
592            preflights: &[PreflightGpu],
593            device_ctx: &openvm_cuda_common::stream::GpuDeviceCtx,
594            required_heights: Option<&[usize]>,
595        ) -> Option<Vec<AirProvingContext<GpuBackend>>> {
596            let blob = StackingBlob::new(child_vk, proofs, preflights, device_ctx);
597            let ctx = (
598                StandardTracegenGpuCtx {
599                    vk: child_vk,
600                    proofs,
601                    preflights,
602                    device_ctx,
603                },
604                &blob,
605            );
606
607            let gpu_chips = [
608                StackingModuleChip::OpeningClaims,
609                StackingModuleChip::StackingClaims,
610            ];
611            let cpu_chips = [
612                StackingModuleChip::UnivariateRound,
613                StackingModuleChip::SumcheckRounds,
614                StackingModuleChip::EqBase,
615                StackingModuleChip::EqBits,
616            ];
617
618            // Launch all GPU tracegen kernels serially first (default stream).
619            let indexed_gpu_proving_ctxs = gpu_chips
620                .iter()
621                .map(|chip| {
622                    (
623                        chip.index(),
624                        chip.generate_proving_ctx(
625                            &ctx,
626                            required_heights.map(|heights| heights[chip.index()]),
627                        ),
628                    )
629                })
630                .collect::<Vec<_>>();
631
632            // Phase 1: CPU trace generation in parallel
633            let proofs_cpu = ctx.0.proofs.iter().map(|p| &p.cpu).collect_vec();
634            let preflights_cpu = ctx.0.preflights.iter().map(|p| &p.cpu).collect_vec();
635            let cpu_ctx = StandardTracegenCtx {
636                vk: &ctx.0.vk.cpu,
637                proofs: &proofs_cpu,
638                preflights: &preflights_cpu,
639            };
640            let span = tracing::Span::current();
641            let indexed_cpu_rm_traces = cpu_chips
642                .par_iter()
643                .map(|chip| {
644                    let _guard = span.enter();
645                    (
646                        chip.index(),
647                        RowMajorChip::generate_trace(
648                            chip,
649                            &cpu_ctx,
650                            required_heights.map(|heights| heights[chip.index()]),
651                        ),
652                    )
653                })
654                .collect::<Vec<_>>();
655
656            // Phase 2: H2D transfer serially on main thread
657            let indexed_cpu_gpu_proving_ctxs = indexed_cpu_rm_traces
658                .into_iter()
659                .map(|(idx, trace)| {
660                    (
661                        idx,
662                        trace.map(|m| {
663                            AirProvingContext::simple_no_pis(
664                                transport_matrix_h2d_row(&m, device_ctx).unwrap(),
665                            )
666                        }),
667                    )
668                })
669                .collect::<Vec<_>>();
670
671            indexed_gpu_proving_ctxs
672                .into_iter()
673                .chain(indexed_cpu_gpu_proving_ctxs)
674                .sorted_by(|a, b| a.0.cmp(&b.0))
675                .map(|(_idx, ctx)| ctx)
676                .collect()
677        }
678    }
679}