openvm_sdk/prover/deferral/
multi_circuit.rs1use std::sync::Arc;
2
3use eyre::Result;
4use itertools::Itertools;
5use openvm_continuations::{
6 prover::{DeferralChildVkKind, DeferralCircuitProver},
7 SC,
8};
9use openvm_deferral_circuit::{DeferralExtension, DeferralFn};
10use openvm_recursion_circuit::utils::poseidon2_hash_slice;
11use openvm_sdk_config::deferral::{DeferralCircuitConfig, DeferralConfig, SupportedDeferral};
12use openvm_stark_backend::{
13 keygen::types::MultiStarkProvingKey,
14 p3_field::{PrimeCharacteristicRing, PrimeField32},
15 proof::Proof,
16 SystemParams,
17};
18use openvm_stark_sdk::config::baby_bear_poseidon2::{
19 poseidon2_compress_with_capacity, DIGEST_SIZE, F,
20};
21use openvm_verify_stark_host::pvs::DeferralPvs;
22use tracing::info_span;
23
24use crate::{
25 config::AggregationConfig,
26 keygen::{AggPrefixProvingKey, DeferralCircuitProvingKey, DeferralProvingKey},
27 prover::SingleDeferralCircuitProver,
28 DeferralInput,
29};
30
31cfg_if::cfg_if! {
32 if #[cfg(feature = "cuda")] {
33 use openvm_continuations::prover::DeferralInnerGpuProver as DeferralInnerProver;
34 use openvm_continuations::prover::DeferralHookGpuProver as DeferralHookProver;
35 type E = openvm_cuda_backend::BabyBearPoseidon2GpuEngine;
36 } else {
37 use openvm_continuations::prover::DeferralInnerCpuProver as DeferralInnerProver;
38 use openvm_continuations::prover::DeferralHookCpuProver as DeferralHookProver;
39 type E = openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2CpuEngine;
40 }
41}
42
43#[allow(clippy::large_enum_variant)]
44#[derive(Clone)]
45pub enum DeferralProof {
46 Present(Proof<SC>),
47 Absent(DeferralPvs<F>),
48}
49
50pub struct MultiDeferralCircuitProver {
52 pub single_circuit_provers: Vec<SingleDeferralCircuitProver>,
53 pub internal_recursive_prover: DeferralInnerProver,
54 pub def_hook_prover: DeferralHookProver,
55}
56
57impl MultiDeferralCircuitProver {
58 pub fn new<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
59 def_circuit_prover: DP,
60 agg_config: AggregationConfig,
61 hook_params: SystemParams,
62 ) -> Self {
63 assert_eq!(def_circuit_prover.get_def_idx(), 0);
64 let single_circuit_prover = SingleDeferralCircuitProver::new(
65 def_circuit_prover,
66 agg_config.params.leaf,
67 agg_config.params.internal.clone(),
68 );
69 let internal_recursive_prover = DeferralInnerProver::new::<E>(
70 single_circuit_prover.internal_for_leaf_prover.get_vk(),
71 agg_config.params.internal,
72 true,
73 );
74 let internal_recursive_cached_commit = internal_recursive_prover
75 .get_vk_commit(true)
76 .cached_commit
77 .into();
78 let def_hook_prover = DeferralHookProver::new::<E>(
79 internal_recursive_prover.get_vk(),
80 internal_recursive_cached_commit,
81 hook_params,
82 );
83 Self {
84 single_circuit_provers: vec![single_circuit_prover],
85 internal_recursive_prover,
86 def_hook_prover,
87 }
88 }
89
90 pub fn from_pks<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
91 def_circuit_prover: DP,
92 def_prefix_pk: AggPrefixProvingKey,
93 def_internal_recursive_pk: Arc<MultiStarkProvingKey<SC>>,
94 def_hook_pk: Arc<MultiStarkProvingKey<SC>>,
95 ) -> Self {
96 assert_eq!(def_circuit_prover.get_def_idx(), 0);
97 let single_circuit_prover = SingleDeferralCircuitProver::from_pks(
98 def_circuit_prover,
99 def_prefix_pk.leaf,
100 def_prefix_pk.internal_for_leaf,
101 );
102 Self::from_single_circuit_prover(
103 single_circuit_prover,
104 def_internal_recursive_pk,
105 def_hook_pk,
106 )
107 }
108
109 pub fn from_single_circuit_prover(
110 single_circuit_prover: SingleDeferralCircuitProver,
111 def_internal_recursive_pk: Arc<MultiStarkProvingKey<SC>>,
112 def_hook_pk: Arc<MultiStarkProvingKey<SC>>,
113 ) -> Self {
114 assert_eq!(single_circuit_prover.def_circuit_prover.get_def_idx(), 0);
115 let internal_recursive_prover = DeferralInnerProver::from_pk::<E>(
116 single_circuit_prover.internal_for_leaf_prover.get_vk(),
117 def_internal_recursive_pk,
118 true,
119 );
120 let internal_recursive_cached_commit = internal_recursive_prover
121 .get_vk_commit(true)
122 .cached_commit
123 .into();
124 let def_hook_prover = DeferralHookProver::from_pk::<E>(
125 internal_recursive_prover.get_vk(),
126 internal_recursive_cached_commit,
127 def_hook_pk,
128 );
129 Self {
130 single_circuit_provers: vec![single_circuit_prover],
131 internal_recursive_prover,
132 def_hook_prover,
133 }
134 }
135
136 pub fn with_prover<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
137 mut self,
138 def_circuit_prover: DP,
139 ) -> Self {
140 assert_eq!(
141 def_circuit_prover.get_def_idx(),
142 self.single_circuit_provers.len()
143 );
144 let leaf_params = self.single_circuit_provers[0]
145 .leaf_prover
146 .get_vk()
147 .inner
148 .params
149 .clone();
150 let internal_params = self.internal_recursive_prover.get_vk().inner.params.clone();
151 let single_circuit_prover =
152 SingleDeferralCircuitProver::new(def_circuit_prover, leaf_params, internal_params);
153 self.single_circuit_provers.push(single_circuit_prover);
154 self
155 }
156
157 pub fn with_prover_from_pk<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
158 self,
159 def_circuit_prover: DP,
160 def_prefix_pk: AggPrefixProvingKey,
161 ) -> Self {
162 assert_eq!(
163 def_circuit_prover.get_def_idx(),
164 self.single_circuit_provers.len()
165 );
166 let single_circuit_prover = SingleDeferralCircuitProver::from_pks(
167 def_circuit_prover,
168 def_prefix_pk.leaf,
169 def_prefix_pk.internal_for_leaf,
170 );
171 self.with_single_circuit_prover(single_circuit_prover)
172 }
173
174 pub fn with_single_circuit_prover(
175 mut self,
176 single_circuit_prover: SingleDeferralCircuitProver,
177 ) -> Self {
178 assert_eq!(
179 single_circuit_prover.def_circuit_prover.get_def_idx(),
180 self.single_circuit_provers.len()
181 );
182 self.single_circuit_provers.push(single_circuit_prover);
183 self
184 }
185
186 pub fn prove(&self, inputs: &[DeferralInput]) -> Result<Vec<DeferralProof>> {
187 let per_circuit = self
189 .single_circuit_provers
190 .iter()
191 .zip_eq(inputs)
192 .map(|(prover, inputs)| prover.prove(inputs))
193 .collect::<Result<Vec<_>>>()?;
194
195 let mut per_circuit = per_circuit
197 .into_iter()
198 .enumerate()
199 .map(|(def_idx, res)| {
200 let mut proofs = res.internal_for_leaf_proofs;
201 if proofs.is_empty() {
202 let def_circuit_commit = self.single_circuit_provers[def_idx]
203 .circuit_commit(self.internal_recursive_prover.get_vk_commit(false));
204 Ok(DeferralProof::Absent(absent_deferral_pvs(
205 def_idx,
206 def_circuit_commit,
207 )))
208 } else {
209 let mut merkle_depth = 2usize;
210 let mut layer = 0usize;
211
212 loop {
216 let is_first = layer == 0;
217 let child_merkle_depth = if proofs.len() > 1 {
218 let d = merkle_depth;
219 merkle_depth += 1;
220 Some(d)
221 } else {
222 None
223 };
224
225 proofs = info_span!(
226 "agg_layer",
227 group = format!("internal_recursive.{layer}"),
228 circuit = def_idx
229 )
230 .in_scope(|| {
231 proofs
232 .chunks(2)
233 .enumerate()
234 .map(|(idx, chunk)| {
235 let kind = if is_first {
236 DeferralChildVkKind::DeferralAggregation
237 } else {
238 DeferralChildVkKind::RecursiveSelf
239 };
240 info_span!("single_internal_agg", idx = idx).in_scope(|| {
241 self.internal_recursive_prover.agg_prove::<E>(
242 chunk,
243 kind,
244 child_merkle_depth,
245 )
246 })
247 })
248 .collect::<Result<Vec<_>>>()
249 })?;
250
251 layer += 1;
252 if proofs.len() == 1 {
253 break;
254 }
255 }
256
257 Ok(DeferralProof::Present(
259 self.def_hook_prover
260 .prove::<E>(proofs.pop().unwrap(), res.leaf_io_commits)?,
261 ))
262 }
263 })
264 .collect::<Result<Vec<_>>>()?;
265
266 let target_length = per_circuit.len().next_power_of_two();
268 for def_idx in per_circuit.len()..target_length {
269 per_circuit.push(DeferralProof::Absent(absent_deferral_pvs(
270 def_idx,
271 [F::ZERO; DIGEST_SIZE],
272 )));
273 }
274
275 Ok(per_circuit)
276 }
277
278 pub fn make_config(&self, supported_deferrals: Vec<SupportedDeferral>) -> DeferralConfig {
279 let vk_commit = self.internal_recursive_prover.get_vk_commit(false);
280 let circuits = self
281 .single_circuit_provers
282 .iter()
283 .zip_eq(supported_deferrals)
284 .map(|(p, deferral)| {
285 let commit = p.circuit_commit(vk_commit).into();
286 DeferralCircuitConfig {
287 def_type: deferral,
288 commit,
289 }
290 })
291 .collect();
292 DeferralConfig::new(circuits)
293 }
294
295 pub fn make_extension(&self, fns: Vec<Arc<DeferralFn>>) -> DeferralExtension {
296 let vk_commit = self.internal_recursive_prover.get_vk_commit(false);
297 let def_circuit_commits = self
298 .single_circuit_provers
299 .iter()
300 .map(|p| {
301 p.circuit_commit(vk_commit)
302 .iter()
303 .flat_map(|f| f.to_unique_u32().to_le_bytes())
304 .collect::<Vec<_>>()
305 .try_into()
306 .unwrap()
307 })
308 .collect();
309 DeferralExtension {
310 fns,
311 def_circuit_commits,
312 }
313 }
314
315 pub fn get_pk(&self) -> DeferralProvingKey {
316 let circuits = self
317 .single_circuit_provers
318 .iter()
319 .map(|single_circuit_prover| DeferralCircuitProvingKey {
320 def_circuit_pk: single_circuit_prover.def_circuit_prover.get_pk(),
321 agg_prefix_pk: AggPrefixProvingKey {
322 leaf: single_circuit_prover.leaf_prover.get_pk(),
323 internal_for_leaf: single_circuit_prover.internal_for_leaf_prover.get_pk(),
324 },
325 })
326 .collect();
327 DeferralProvingKey {
328 circuits,
329 def_internal_recursive_pk: self.internal_recursive_prover.get_pk(),
330 def_hook_pk: self.def_hook_prover.get_pk(),
331 }
332 }
333}
334
335fn absent_deferral_pvs(def_idx: usize, def_circuit_commit: [F; DIGEST_SIZE]) -> DeferralPvs<F> {
336 let input_acc_hash = poseidon2_hash_slice(&def_circuit_commit).0;
337 let output_acc_hash = poseidon2_hash_slice(&[F::ZERO]).0;
338 let combined_hash = poseidon2_compress_with_capacity(input_acc_hash, output_acc_hash).0;
339 DeferralPvs {
340 initial_acc_hash: combined_hash,
341 final_acc_hash: combined_hash,
342 depth: F::ONE,
343 node_idx: F::from_usize(def_idx),
344 }
345}