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
35const 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 #[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 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; 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 if !cur_is_sample {
88 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 if cur_is_sample {
102 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 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 break;
173 }
174
175 if preflight.transcript.samples()[tidx] != is_sample {
176 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 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 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 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 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}