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 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 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 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 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 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 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 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}