openvm_recursion_circuit/transcript/
mod.rs

1use core::borrow::BorrowMut;
2use std::sync::Arc;
3
4use itertools::Itertools;
5use openvm_cpu_backend::CpuBackend;
6use openvm_poseidon2_air::{Poseidon2Config, Poseidon2SubChip, POSEIDON2_WIDTH};
7use openvm_stark_backend::{
8    keygen::types::MultiStarkVerifyingKey, p3_maybe_rayon::prelude::*, proof::Proof,
9    prover::AirProvingContext, AirRef, StarkProtocolConfig, SystemParams,
10};
11use openvm_stark_sdk::config::baby_bear_poseidon2::{poseidon2_perm, BabyBearPoseidon2Config, F};
12use p3_air::BaseAir;
13use p3_baby_bear::Poseidon2BabyBear;
14use p3_field::{PrimeCharacteristicRing, PrimeField32};
15use p3_matrix::dense::RowMajorMatrix;
16use p3_symmetric::Permutation;
17use tracing::trace_span;
18
19use crate::{
20    system::{AirModule, BusInventory, GlobalCtxCpu, Preflight, TraceGenModule},
21    transcript::{
22        merkle_verify::{MerkleVerifyAir, MerkleVerifyCols},
23        poseidon2::{Poseidon2Air, Poseidon2Cols, CHUNK},
24        transcript::{TranscriptAir, TranscriptCols},
25    },
26};
27
28#[cfg(feature = "cuda")]
29mod cuda_abi;
30pub mod merkle_verify;
31pub mod poseidon2;
32#[allow(clippy::module_inception)]
33pub mod transcript;
34
35// Should be 1 when 3 <= max_constraint_degree < 7
36const SBOX_REGISTERS: usize = 1;
37
38pub struct TranscriptModule {
39    pub bus_inventory: BusInventory,
40    params: SystemParams,
41    final_state_bus_enabled: bool,
42
43    sub_chip: Poseidon2SubChip<F, SBOX_REGISTERS>,
44    perm: Poseidon2BabyBear<POSEIDON2_WIDTH>,
45}
46
47impl TranscriptModule {
48    pub fn new(
49        bus_inventory: BusInventory,
50        params: SystemParams,
51        final_state_bus_enabled: bool,
52    ) -> Self {
53        let sub_chip = Poseidon2SubChip::<F, 1>::new(Poseidon2Config::default().constants);
54        Self {
55            bus_inventory,
56            params,
57            final_state_bus_enabled,
58            sub_chip,
59            perm: poseidon2_perm().clone(),
60        }
61    }
62
63    // Builds trace for transcript and merkle verify AIRs (and records poseidon2 permutations).
64    // Also combines in the poseidon2 permutations from preflight (from WHIR).
65    #[tracing::instrument(name = "generate_trace", level = "trace", skip_all)]
66    fn build_trace_artifacts(
67        &self,
68        preflights: &[Preflight],
69        mut poseidon2_perm_inputs: Vec<[F; POSEIDON2_WIDTH]>,
70        mut poseidon2_compress_inputs: Vec<[F; POSEIDON2_WIDTH]>,
71        required_height: Option<usize>,
72    ) -> Option<TranscriptTraceArtifacts> {
73        let transcript_width = TranscriptCols::<F>::width();
74        let mut valid_rows = Vec::with_capacity(preflights.len());
75
76        let mut transcript_valid_rows = 0;
77        // First pass, calculate number of rows for transcript
78        for preflight in preflights.iter() {
79            poseidon2_perm_inputs.extend_from_slice(&preflight.poseidon2_perm_inputs);
80            poseidon2_compress_inputs.extend_from_slice(&preflight.poseidon2_compress_inputs);
81            let mut cur_is_sample = false; // should always start with observe?
82            let mut count = 0;
83            let mut num_valid_rows: usize = 0;
84            for op_is_sample in preflight.transcript.samples() {
85                if *op_is_sample {
86                    // sample
87                    if !cur_is_sample {
88                        // observe -> sample, need a new row and permute
89                        num_valid_rows += 1;
90                        cur_is_sample = true;
91                        count = 1;
92                    } else {
93                        if count == CHUNK {
94                            num_valid_rows += 1;
95                            count = 0;
96                        }
97                        count += 1;
98                    }
99                } else {
100                    // observe
101                    if cur_is_sample {
102                        // sample -> observe, no need to permute, but still need a new row
103                        num_valid_rows += 1;
104                        cur_is_sample = false;
105                        count = 1;
106                    } else {
107                        if count == CHUNK {
108                            num_valid_rows += 1;
109                            count = 0;
110                        }
111                        count += 1;
112                    }
113                }
114            }
115            if count > 0 {
116                num_valid_rows += 1;
117            }
118            valid_rows.push(num_valid_rows);
119            transcript_valid_rows += num_valid_rows;
120        }
121        let transcript_num_rows = if let Some(height) = required_height {
122            if height < transcript_valid_rows {
123                return None;
124            }
125            height
126        } else {
127            transcript_valid_rows.next_power_of_two()
128        };
129        let mut transcript_trace = vec![F::ZERO; transcript_num_rows * transcript_width];
130
131        let mut skip = 0;
132        // Second pass, fill in the transcript trace.
133        for (pidx, preflight) in preflights.iter().enumerate() {
134            let mut tidx = 0;
135            let mut prev_poseidon_state = [F::ZERO; POSEIDON2_WIDTH];
136            let off = skip * transcript_width;
137            let end = off + valid_rows[pidx] * transcript_width;
138            for (i, row) in transcript_trace[off..end]
139                .chunks_exact_mut(transcript_width)
140                .enumerate()
141            {
142                let cols: &mut TranscriptCols<F> = row.borrow_mut();
143                cols.proof_idx = F::from_usize(pidx);
144                if i == 0 {
145                    cols.is_proof_start = F::ONE;
146                }
147                let is_sample = preflight.transcript.samples()[tidx];
148
149                cols.is_sample = F::from_bool(is_sample);
150                cols.tidx = F::from_usize(tidx);
151                cols.mask[0] = F::from_bool(true);
152
153                cols.prev_state = prev_poseidon_state;
154
155                if is_sample {
156                    debug_assert_eq!(
157                        cols.prev_state[CHUNK - 1],
158                        preflight.transcript.values()[tidx],
159                        "sample value mismatch",
160                    );
161                } else {
162                    cols.prev_state[0] = preflight.transcript.values()[tidx];
163                }
164
165                tidx += 1;
166                let mut idx: usize = 1;
167
168                let mut permuted = false;
169                loop {
170                    if tidx >= preflight.transcript.len() {
171                        // at the end, no permutation needed
172                        break;
173                    }
174
175                    if preflight.transcript.samples()[tidx] != is_sample {
176                        // encounter a different type of operation. Permute if it's going to sample
177                        permuted = preflight.transcript.samples()[tidx];
178                        break;
179                    }
180
181                    cols.mask[idx] = F::from_bool(true);
182                    if is_sample {
183                        debug_assert_eq!(
184                            cols.prev_state[CHUNK - 1 - idx],
185                            preflight.transcript.values()[tidx],
186                            "sample value mismatch",
187                        );
188                    } else {
189                        cols.prev_state[idx] = preflight.transcript.values()[tidx];
190                    }
191
192                    tidx += 1;
193                    idx += 1;
194                    if idx == CHUNK {
195                        // If it's sample -> observe, we don't need to permute. otherwise permute
196                        permuted = tidx < preflight.transcript.len()
197                            && (!is_sample || preflight.transcript.samples()[tidx]);
198                        break;
199                    }
200                }
201
202                prev_poseidon_state = cols.prev_state;
203                if permuted {
204                    self.perm.permute_mut(&mut prev_poseidon_state);
205                    poseidon2_perm_inputs.push(cols.prev_state);
206                }
207                cols.post_state = prev_poseidon_state;
208            }
209            skip += valid_rows[pidx];
210            assert_eq!(tidx, preflight.transcript.len());
211        }
212
213        Some(TranscriptTraceArtifacts {
214            transcript_trace: RowMajorMatrix::new(transcript_trace, transcript_width),
215            poseidon2_perm_inputs,
216            poseidon2_compress_inputs,
217        })
218    }
219
220    fn dedup_poseidon_inputs(
221        poseidon2_perm_inputs: Vec<[F; POSEIDON2_WIDTH]>,
222        poseidon2_compress_inputs: Vec<[F; POSEIDON2_WIDTH]>,
223    ) -> (Vec<[F; POSEIDON2_WIDTH]>, Vec<Poseidon2Count>) {
224        let keyed_perm_states = poseidon2_perm_inputs
225            .into_iter()
226            .map(|state| (state.map(|x| x.as_canonical_u32()), state, true));
227        let keyed_compress_states = poseidon2_compress_inputs
228            .into_iter()
229            .map(|state| (state.map(|x| x.as_canonical_u32()), state, false));
230        let mut keyed_states = keyed_perm_states
231            .into_iter()
232            .chain(keyed_compress_states)
233            .collect_vec();
234        keyed_states.sort_unstable_by(|a, b| a.0.cmp(&b.0));
235
236        let mut deduped = Vec::new();
237        let mut counts: Vec<Poseidon2Count> = Vec::new();
238        let mut last_key: Option<[u32; POSEIDON2_WIDTH]> = None;
239
240        for (key, state, is_perm) in keyed_states {
241            if last_key == Some(key) {
242                if is_perm {
243                    counts.last_mut().unwrap().perm += 1;
244                } else {
245                    counts.last_mut().unwrap().compress += 1;
246                }
247            } else {
248                deduped.push(state);
249                counts.push(if is_perm {
250                    Poseidon2Count {
251                        perm: 1,
252                        compress: 0,
253                    }
254                } else {
255                    Poseidon2Count {
256                        perm: 0,
257                        compress: 1,
258                    }
259                });
260                last_key = Some(key);
261            }
262        }
263        (deduped, counts)
264    }
265}
266
267impl AirModule for TranscriptModule {
268    fn num_airs(&self) -> usize {
269        3
270    }
271
272    fn airs<SC: StarkProtocolConfig<F = F>>(&self) -> Vec<AirRef<SC>> {
273        let transcript_air = TranscriptAir {
274            transcript_bus: self.bus_inventory.transcript_bus,
275            poseidon2_permute_bus: self.bus_inventory.poseidon2_permute_bus,
276            final_state_bus: self
277                .final_state_bus_enabled
278                .then_some(self.bus_inventory.final_state_bus),
279        };
280        let poseidon2_air = Poseidon2Air::<F, SBOX_REGISTERS> {
281            subair: self.sub_chip.air.clone(),
282            poseidon2_permute_bus: self.bus_inventory.poseidon2_permute_bus,
283            poseidon2_compress_bus: self.bus_inventory.poseidon2_compress_bus,
284        };
285        let merkle_verify_air = MerkleVerifyAir {
286            poseidon2_compress_bus: self.bus_inventory.poseidon2_compress_bus,
287            merkle_verify_bus: self.bus_inventory.merkle_verify_bus,
288            commitments_bus: self.bus_inventory.commitments_bus,
289            right_shift_bus: self.bus_inventory.right_shift_bus,
290            k: self.params.k_whir(),
291        };
292        vec![
293            Arc::new(transcript_air),
294            Arc::new(poseidon2_air),
295            Arc::new(merkle_verify_air),
296        ]
297    }
298}
299
300pub(super) struct TranscriptTraceArtifacts {
301    transcript_trace: RowMajorMatrix<F>,
302    poseidon2_perm_inputs: Vec<[F; POSEIDON2_WIDTH]>,
303    poseidon2_compress_inputs: Vec<[F; POSEIDON2_WIDTH]>,
304}
305
306#[repr(C)]
307#[derive(Copy, Clone, Default)]
308pub(super) struct Poseidon2Count {
309    pub perm: u32,
310    pub compress: u32,
311}
312
313impl<SC: StarkProtocolConfig<F = F>> TraceGenModule<GlobalCtxCpu, CpuBackend<SC>>
314    for TranscriptModule
315{
316    // External Poseidon2 compress inputs
317    type ModuleSpecificCtx<'a> = (&'a Vec<[F; POSEIDON2_WIDTH]>, &'a Vec<[F; POSEIDON2_WIDTH]>);
318
319    #[tracing::instrument(skip_all)]
320    fn generate_proving_ctxs(
321        &self,
322        child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
323        proofs: &[Proof<BabyBearPoseidon2Config>],
324        preflights: &[Preflight],
325        ctx: &Self::ModuleSpecificCtx<'_>,
326        required_heights: Option<&[usize]>,
327    ) -> Option<Vec<AirProvingContext<CpuBackend<SC>>>> {
328        let external_poseidon2_permute_inputs = ctx.0;
329        let external_poseidon2_compress_inputs = ctx.1;
330        let (required_transcript, required_poseidon2, required_merkle_verify) =
331            if let Some(heights) = required_heights {
332                if heights.len() != 3 {
333                    return None;
334                }
335                (Some(heights[0]), Some(heights[1]), Some(heights[2]))
336            } else {
337                (None, None, None)
338            };
339
340        let (merkle_verify_trace_vec, poseidon2_compress_inputs) =
341            tracing::info_span!("wrapper.generate_trace", air = "MerkleVerify").in_scope(|| {
342                merkle_verify::generate_trace(
343                    child_vk,
344                    proofs,
345                    preflights,
346                    &self.params,
347                    required_merkle_verify,
348                )
349            })?;
350        let merkle_verify_trace =
351            RowMajorMatrix::new(merkle_verify_trace_vec, MerkleVerifyCols::<F>::width());
352        let TranscriptTraceArtifacts {
353            transcript_trace,
354            mut poseidon2_perm_inputs,
355            mut poseidon2_compress_inputs,
356        } = tracing::trace_span!("wrapper.generate_trace", air = "Transcript").in_scope(|| {
357            self.build_trace_artifacts(
358                preflights,
359                vec![],
360                poseidon2_compress_inputs,
361                required_transcript,
362            )
363        })?;
364        poseidon2_perm_inputs.extend_from_slice(external_poseidon2_permute_inputs);
365        poseidon2_compress_inputs.extend_from_slice(external_poseidon2_compress_inputs);
366
367        let poseidon2_trace =
368            trace_span!("wrapper.generate_trace", air = "Poseidon2").in_scope(|| {
369                trace_span!("generate_trace").in_scope(|| {
370                    let (mut poseidon_states, poseidon_counts) = Self::dedup_poseidon_inputs(
371                        poseidon2_perm_inputs,
372                        poseidon2_compress_inputs,
373                    );
374                    let poseidon2_valid_rows = poseidon_states.len();
375                    let poseidon2_num_rows = if let Some(height) = required_poseidon2 {
376                        if height == 0 || poseidon2_valid_rows > height {
377                            return None;
378                        }
379                        height
380                    } else if poseidon2_valid_rows == 0 {
381                        1
382                    } else {
383                        poseidon2_valid_rows.next_power_of_two()
384                    };
385                    poseidon_states.resize(poseidon2_num_rows, [F::ZERO; POSEIDON2_WIDTH]);
386
387                    let inner_width = self.sub_chip.air.width();
388                    let poseidon2_width = Poseidon2Cols::<F, SBOX_REGISTERS>::width();
389                    let inner_trace = self.sub_chip.generate_trace(poseidon_states);
390                    let mut poseidon_trace = F::zero_vec(poseidon2_num_rows * poseidon2_width);
391
392                    poseidon_trace
393                        .par_chunks_mut(poseidon2_width)
394                        .zip(inner_trace.values.par_chunks(inner_width))
395                        .enumerate()
396                        .for_each(|(i, (row, inner_row))| {
397                            row[..inner_width].copy_from_slice(inner_row);
398                            let cols: &mut Poseidon2Cols<F, SBOX_REGISTERS> = row.borrow_mut();
399                            let count = poseidon_counts.get(i).copied().unwrap_or_default();
400                            cols.permute_mult = F::from_u32(count.perm);
401                            cols.compress_mult = F::from_u32(count.compress);
402                        });
403                    Some(RowMajorMatrix::new(poseidon_trace, poseidon2_width))
404                })
405            })?;
406
407        // Finally, make the RawInput structs
408        Some(
409            [transcript_trace, poseidon2_trace, merkle_verify_trace]
410                .map(AirProvingContext::simple_no_pis)
411                .into_iter()
412                .collect(),
413        )
414    }
415}
416
417#[cfg(feature = "cuda")]
418mod cuda_tracegen {
419    use itertools::Itertools;
420    use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
421    use openvm_cuda_common::{
422        copy::{MemCopyD2H, MemCopyH2D},
423        d_buffer::DeviceBuffer,
424        stream::GpuDeviceCtx,
425    };
426    use openvm_stark_backend::prover::MatrixDimensions;
427
428    use super::*;
429    use crate::{
430        cuda::{preflight::PreflightGpu, proof::ProofGpu, vk::VerifyingKeyGpu, GlobalCtxGpu},
431        transcript::{
432            cuda_abi,
433            merkle_verify::{self, cuda::MerkleVerifyBlob},
434            transcript::cuda::TranscriptAirBlob,
435        },
436    };
437
438    pub(crate) struct TranscriptBlob {
439        pub merkle_verify_blob: MerkleVerifyBlob,
440        pub transcript_air_blob: TranscriptAirBlob,
441
442        // Because we currently can only copy to the beginning of a DeviceBuffer, the layout is
443        // expected to be in this order:
444        // - Preflight permutations
445        // - Preflight compressions
446        // - Merkle verify compressions
447        // - Transcript permutations
448        pub poseidon2_buffer: DeviceBuffer<F>,
449        pub num_prefix_perms: usize,
450        pub num_suffix_perms: usize,
451        pub num_compress_inputs: usize,
452    }
453
454    impl TranscriptBlob {
455        #[tracing::instrument(name = "generate_blob", skip_all)]
456        pub fn new(
457            child_vk: &VerifyingKeyGpu,
458            proofs: &[ProofGpu],
459            preflights: &[PreflightGpu],
460            external_poseidon2_inputs: &(
461                &Vec<[F; POSEIDON2_WIDTH]>,
462                &Vec<[F; POSEIDON2_WIDTH]>,
463                &GpuDeviceCtx,
464            ),
465        ) -> Self {
466            let external_poseidon2_permute_inputs = external_poseidon2_inputs.0;
467            let external_poseidon2_compress_inputs = external_poseidon2_inputs.1;
468            let device_ctx = external_poseidon2_inputs.2;
469            let poseidon2_perm_inputs = preflights
470                .iter()
471                .flat_map(|preflight| preflight.cpu.poseidon2_perm_inputs.clone())
472                .chain(external_poseidon2_permute_inputs.iter().copied())
473                .collect_vec();
474            let poseidon2_compress_inputs = preflights
475                .iter()
476                .flat_map(|preflight| preflight.cpu.poseidon2_compress_inputs.clone())
477                .chain(external_poseidon2_compress_inputs.iter().copied())
478                .collect_vec();
479            let num_prefix_perms = poseidon2_perm_inputs.len();
480            let mut num_compress_inputs = poseidon2_compress_inputs.len();
481
482            let merkle_verify_blob = MerkleVerifyBlob::new(
483                child_vk,
484                proofs,
485                preflights,
486                num_prefix_perms + num_compress_inputs,
487            );
488            num_compress_inputs += merkle_verify_blob.total_rows;
489
490            let transcript_air_blob =
491                TranscriptAirBlob::new(preflights, (num_prefix_perms + num_compress_inputs) as u32);
492            let num_suffix_perms = transcript_air_blob.num_poseidon2_perms;
493
494            let mut poseidon2_buffer = DeviceBuffer::with_capacity_on(
495                (num_prefix_perms + num_compress_inputs + num_suffix_perms) * POSEIDON2_WIDTH,
496                device_ctx,
497            );
498            poseidon2_perm_inputs
499                .into_iter()
500                .flatten()
501                .chain(poseidon2_compress_inputs.into_iter().flatten())
502                .collect_vec()
503                .copy_to_on(&mut poseidon2_buffer, device_ctx)
504                .unwrap();
505
506            Self {
507                merkle_verify_blob,
508                transcript_air_blob,
509                poseidon2_buffer,
510                num_prefix_perms,
511                num_suffix_perms,
512                num_compress_inputs,
513            }
514        }
515    }
516
517    impl TraceGenModule<GlobalCtxGpu, GpuBackend> for TranscriptModule {
518        type ModuleSpecificCtx<'a> = (
519            &'a Vec<[F; POSEIDON2_WIDTH]>,
520            &'a Vec<[F; POSEIDON2_WIDTH]>,
521            &'a openvm_cuda_common::stream::GpuDeviceCtx,
522        );
523
524        #[tracing::instrument(skip_all)]
525        fn generate_proving_ctxs(
526            &self,
527            child_vk: &VerifyingKeyGpu,
528            proofs: &[ProofGpu],
529            preflights: &[PreflightGpu],
530            ctx: &Self::ModuleSpecificCtx<'_>,
531            required_heights: Option<&[usize]>,
532        ) -> Option<Vec<AirProvingContext<GpuBackend>>> {
533            let device_ctx = ctx.2;
534            let (required_transcript, required_poseidon2, required_merkle_verify) =
535                if let Some(heights) = required_heights {
536                    if heights.len() != 3 {
537                        return None;
538                    }
539                    (Some(heights[0]), Some(heights[1]), Some(heights[2]))
540                } else {
541                    (None, None, None)
542                };
543            let blob = TranscriptBlob::new(child_vk, proofs, preflights, ctx);
544
545            let merkle_trace = tracing::trace_span!("wrapper.generate_trace", air = "MerkleVerify")
546                .in_scope(|| {
547                    merkle_verify::cuda::generate_trace(&blob, device_ctx, required_merkle_verify)
548                })?;
549            let transcript_trace = tracing::trace_span!(
550                "wrapper.generate_trace",
551                air = "Transcript"
552            )
553            .in_scope(|| {
554                transcript::cuda::generate_trace(preflights, &blob, device_ctx, required_transcript)
555            })?;
556            let poseidon_trace = trace_span!("wrapper.generate_trace", air = "Poseidon2")
557                .in_scope(|| {
558                    trace_span!("generate_trace").in_scope(|| {
559                        let poseidon2_width = Poseidon2Cols::<F, SBOX_REGISTERS>::width();
560                        let total_poseidon2_inputs = blob.num_prefix_perms
561                            + blob.num_compress_inputs
562                            + blob.num_suffix_perms;
563
564                        let d_counts = if total_poseidon2_inputs == 0 {
565                            DeviceBuffer::<Poseidon2Count>::new()
566                        } else {
567                            DeviceBuffer::<Poseidon2Count>::with_capacity_on(
568                                total_poseidon2_inputs,
569                                device_ctx,
570                            )
571                        };
572                        let d_records_dedup = if total_poseidon2_inputs == 0 {
573                            DeviceBuffer::<F>::new()
574                        } else {
575                            DeviceBuffer::<F>::with_capacity_on(
576                                total_poseidon2_inputs * POSEIDON2_WIDTH,
577                                device_ctx,
578                            )
579                        };
580                        let d_counts_dedup = if total_poseidon2_inputs == 0 {
581                            DeviceBuffer::<Poseidon2Count>::new()
582                        } else {
583                            DeviceBuffer::<Poseidon2Count>::with_capacity_on(
584                                total_poseidon2_inputs,
585                                device_ctx,
586                            )
587                        };
588
589                        let mut num_records = total_poseidon2_inputs;
590                        if num_records > 0 {
591                            unsafe {
592                                let d_num_records = [num_records].to_device_on(device_ctx).unwrap();
593                                let mut temp_bytes = 0;
594                                cuda_abi::poseidon2_deduplicate_records_get_temp_bytes(
595                                    &blob.poseidon2_buffer,
596                                    &d_counts,
597                                    num_records,
598                                    &d_num_records,
599                                    &mut temp_bytes,
600                                    device_ctx.stream.as_raw(),
601                                )
602                                .unwrap();
603                                let d_temp_storage = if temp_bytes == 0 {
604                                    DeviceBuffer::<u8>::new()
605                                } else {
606                                    DeviceBuffer::<u8>::with_capacity_on(temp_bytes, device_ctx)
607                                };
608                                cuda_abi::poseidon2_deduplicate_records(
609                                    &blob.poseidon2_buffer,
610                                    &d_counts,
611                                    &d_records_dedup,
612                                    &d_counts_dedup,
613                                    num_records,
614                                    &d_num_records,
615                                    blob.num_prefix_perms,
616                                    blob.num_compress_inputs,
617                                    blob.num_suffix_perms,
618                                    &d_temp_storage,
619                                    temp_bytes,
620                                    device_ctx.stream.as_raw(),
621                                )
622                                .unwrap();
623                                num_records = *d_num_records
624                                    .to_host_on(device_ctx)
625                                    .unwrap()
626                                    .first()
627                                    .unwrap();
628                            }
629                        }
630                        let poseidon2_num_rows = if let Some(height) = required_poseidon2 {
631                            if height < num_records {
632                                return None;
633                            }
634                            height
635                        } else if num_records == 0 {
636                            1
637                        } else {
638                            num_records.next_power_of_two()
639                        };
640                        let poseidon_trace_gpu = DeviceMatrix::<F>::with_capacity_on(
641                            poseidon2_num_rows,
642                            poseidon2_width,
643                            device_ctx,
644                        );
645                        unsafe {
646                            cuda_abi::poseidon2_tracegen(
647                                poseidon_trace_gpu.buffer(),
648                                poseidon_trace_gpu.height(),
649                                poseidon_trace_gpu.width(),
650                                &d_records_dedup,
651                                &d_counts_dedup,
652                                num_records,
653                                SBOX_REGISTERS,
654                                device_ctx.stream.as_raw(),
655                            )
656                            .unwrap();
657                        }
658                        Some(poseidon_trace_gpu)
659                    })
660                })?;
661
662            Some(vec![
663                AirProvingContext::simple_no_pis(transcript_trace),
664                AirProvingContext::simple_no_pis(poseidon_trace),
665                AirProvingContext::simple_no_pis(merkle_trace),
666            ])
667        }
668    }
669}