openvm_continuations/circuit/root/
trace.rs1use std::borrow::Borrow;
2
3use itertools::Itertools;
4use openvm_circuit::{
5 arch::POSEIDON2_WIDTH,
6 system::memory::{dimensions::MemoryDimensions, merkle::public_values::UserPublicValuesProof},
7};
8#[cfg(feature = "cuda")]
9use openvm_circuit_primitives::hybrid_chip::cpu_proving_ctx_to_gpu;
10use openvm_cpu_backend::CpuBackend;
11#[cfg(feature = "cuda")]
12use openvm_cuda_backend::{BabyBearBn254Poseidon2HashScheme, GenericGpuBackend};
13#[cfg(feature = "cuda")]
14use openvm_cuda_common::stream::GpuDeviceCtx;
15use openvm_stark_backend::{
16 proof::Proof,
17 prover::{AirProvingContext, ProverBackend},
18 StarkProtocolConfig,
19};
20use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, DIGEST_SIZE, F};
21use openvm_verify_stark_host::pvs::{DeferralPvs, DEF_PVS_AIR_ID};
22use p3_field::PrimeField32;
23
24use crate::circuit::{
25 deferral::DeferralMerkleProofs,
26 root::{commit, memory},
27 SingleAirTraceData, SubCircuitTraceData,
28};
29
30pub trait RootTraceGen<PB: ProverBackend, DC: Clone + Send + Sync> {
33 fn new(deferral_enabled: bool) -> Self;
34 fn generate_pre_verifier_subcircuit_ctx(
35 &self,
36 proof: &Proof<BabyBearPoseidon2Config>,
37 user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, PB::Val>,
38 memory_dimensions: MemoryDimensions,
39 device_ctx: &DC,
40 ) -> SubCircuitTraceData<PB>;
41 fn generate_other_proving_ctxs(
42 &self,
43 proof: &Proof<BabyBearPoseidon2Config>,
44 memory_dimensions: MemoryDimensions,
45 deferral_merkle_proofs: Option<&DeferralMerkleProofs<PB::Val>>,
46 device_ctx: &DC,
47 ) -> (Vec<AirProvingContext<PB>>, Vec<[PB::Val; POSEIDON2_WIDTH]>);
48}
49
50pub struct RootTraceGenImpl {
51 pub deferral_enabled: bool,
52}
53
54impl<SC: StarkProtocolConfig<F = F>> RootTraceGen<CpuBackend<SC>, ()> for RootTraceGenImpl {
55 fn new(deferral_enabled: bool) -> Self {
56 Self { deferral_enabled }
57 }
58
59 fn generate_pre_verifier_subcircuit_ctx(
60 &self,
61 proof: &Proof<BabyBearPoseidon2Config>,
62 user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, F>,
63 memory_dimensions: MemoryDimensions,
64 _device_ctx: &(),
65 ) -> SubCircuitTraceData<CpuBackend<SC>> {
66 let SingleAirTraceData {
67 air_proving_ctx: verifier_ctx,
68 poseidon2_compress_inputs: verifier_compress_inputs,
69 poseidon2_permute_inputs: verifier_permute_inputs,
70 range_check_inputs,
71 } = super::verifier::generate_proving_ctx(proof, self.deferral_enabled);
72 let (commit_ctx, commit_inputs) =
73 commit::generate_proving_ctx(user_pvs_proof.public_values.clone());
74 let (memory_ctx, memory_inputs) = memory::generate_proving_input(
75 user_pvs_proof.public_values_commit,
76 &user_pvs_proof.proof,
77 memory_dimensions,
78 user_pvs_proof.public_values.len(),
79 );
80 SubCircuitTraceData {
81 air_proving_ctxs: vec![verifier_ctx, commit_ctx, memory_ctx],
82 poseidon2_compress_inputs: verifier_compress_inputs
83 .into_iter()
84 .chain(commit_inputs)
85 .chain(memory_inputs)
86 .collect_vec(),
87 poseidon2_permute_inputs: verifier_permute_inputs,
88 range_check_inputs,
89 }
90 }
91
92 fn generate_other_proving_ctxs(
93 &self,
94 proof: &Proof<BabyBearPoseidon2Config>,
95 memory_dimensions: MemoryDimensions,
96 deferral_merkle_proofs: Option<&DeferralMerkleProofs<F>>,
97 _device_ctx: &(),
98 ) -> (
99 Vec<AirProvingContext<CpuBackend<SC>>>,
100 Vec<[F; POSEIDON2_WIDTH]>,
101 ) {
102 let (paths_ctx, paths_inputs) = if let Some(deferral_merkle_proofs) = deferral_merkle_proofs
103 {
104 assert!(self.deferral_enabled);
105 let def_pvs: &DeferralPvs<F> = proof.public_values[DEF_PVS_AIR_ID].as_slice().borrow();
106 let depth = def_pvs.depth.as_canonical_u32() as usize;
107 let (ctx, inputs) = super::def_paths::generate_proving_input(
108 def_pvs.initial_acc_hash,
109 def_pvs.final_acc_hash,
110 &deferral_merkle_proofs.initial_merkle_proof,
111 &deferral_merkle_proofs.final_merkle_proof,
112 memory_dimensions,
113 depth,
114 depth == 0,
115 );
116 (Some(ctx), inputs)
117 } else {
118 assert!(!self.deferral_enabled);
119 (None, vec![])
120 };
121 (paths_ctx.into_iter().collect_vec(), paths_inputs)
122 }
123}
124
125#[cfg(feature = "cuda")]
126impl RootTraceGen<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>, GpuDeviceCtx>
127 for RootTraceGenImpl
128{
129 fn new(deferral_enabled: bool) -> Self {
130 Self { deferral_enabled }
131 }
132
133 fn generate_pre_verifier_subcircuit_ctx(
134 &self,
135 proof: &Proof<BabyBearPoseidon2Config>,
136 user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, F>,
137 memory_dimensions: MemoryDimensions,
138 device_ctx: &GpuDeviceCtx,
139 ) -> SubCircuitTraceData<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>> {
140 let data: SubCircuitTraceData<CpuBackend<BabyBearPoseidon2Config>> =
141 <Self as RootTraceGen<CpuBackend<BabyBearPoseidon2Config>, ()>>::generate_pre_verifier_subcircuit_ctx(
142 self,
143 proof,
144 user_pvs_proof,
145 memory_dimensions,
146 &(),
147 );
148 SubCircuitTraceData {
149 air_proving_ctxs: data
150 .air_proving_ctxs
151 .into_iter()
152 .map(|c| cpu_proving_ctx_to_gpu::<BabyBearBn254Poseidon2HashScheme>(c, device_ctx))
153 .collect_vec(),
154 poseidon2_compress_inputs: data.poseidon2_compress_inputs,
155 poseidon2_permute_inputs: data.poseidon2_permute_inputs,
156 range_check_inputs: data.range_check_inputs,
157 }
158 }
159
160 fn generate_other_proving_ctxs(
161 &self,
162 proof: &Proof<BabyBearPoseidon2Config>,
163 memory_dimensions: MemoryDimensions,
164 deferral_merkle_proofs: Option<&DeferralMerkleProofs<F>>,
165 device_ctx: &GpuDeviceCtx,
166 ) -> (
167 Vec<AirProvingContext<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>>>,
168 Vec<[F; POSEIDON2_WIDTH]>,
169 ) {
170 let (cpu_ctxs, inputs) =
171 <Self as RootTraceGen<CpuBackend<BabyBearPoseidon2Config>, ()>>::generate_other_proving_ctxs(
172 self,
173 proof,
174 memory_dimensions,
175 deferral_merkle_proofs,
176 &(),
177 );
178 let gpu_proving_ctxs = cpu_ctxs
179 .into_iter()
180 .map(|c| cpu_proving_ctx_to_gpu::<BabyBearBn254Poseidon2HashScheme>(c, device_ctx))
181 .collect_vec();
182 (gpu_proving_ctxs, inputs)
183 }
184}