1use core::cmp::Reverse;
2use std::sync::Arc;
3
4use itertools::{izip, Itertools};
5use openvm_circuit_primitives::encoder::Encoder;
6use openvm_cpu_backend::CpuBackend;
7use openvm_stark_backend::{
8 keygen::types::{MultiStarkVerifyingKey, VerifierSinglePreprocessedData},
9 proof::Proof,
10 prover::AirProvingContext,
11 AirRef, FiatShamirTranscript, StarkProtocolConfig, TranscriptHistory,
12};
13use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, Digest, F};
14use p3_field::PrimeCharacteristicRing;
15use p3_matrix::dense::RowMajorMatrix;
16use p3_maybe_rayon::prelude::{IntoParallelRefIterator, ParallelIterator};
17
18use crate::{
19 primitives::{
20 bus::{PowerCheckerBus, RangeCheckerBus},
21 pow::PowerCheckerCpuTraceGenerator,
22 range::{RangeCheckerAir, RangeCheckerCpuTraceGenerator},
23 },
24 proof_shape::{
25 bus::{NumPublicValuesBus, ProofShapePermutationBus, StartingTidxBus},
26 proof_shape::ProofShapeAir,
27 pvs::PublicValuesAir,
28 },
29 system::{
30 frame::MultiStarkVkeyFrame, AirModule, BusIndexManager, BusInventory, GlobalCtxCpu,
31 Preflight, ProofShapePreflight, TraceGenModule, POW_CHECKER_HEIGHT,
32 },
33 tracegen::{ModuleChip, RowMajorChip},
34};
35
36pub mod bus;
37#[allow(clippy::module_inception)]
38pub mod proof_shape;
39pub mod pvs;
40
41#[cfg(feature = "cuda")]
42mod cuda_abi;
43
44#[derive(Clone)]
45pub struct AirMetadata {
46 is_required: bool,
47 need_rot: bool,
48 num_public_values: usize,
49 num_interactions: usize,
50 main_width: usize,
51 cached_widths: Vec<usize>,
52 preprocessed_width: Option<usize>,
53 preprocessed_data: Option<VerifierSinglePreprocessedData<Digest>>,
54}
55
56pub struct ProofShapeModule {
57 per_air: Vec<AirMetadata>,
59 l_skip: usize,
60 max_interaction_count: u32,
64
65 bus_inventory: BusInventory,
67 range_bus: RangeCheckerBus,
68 pow_bus: PowerCheckerBus,
69 permutation_bus: ProofShapePermutationBus,
70 starting_tidx_bus: StartingTidxBus,
71 num_pvs_bus: NumPublicValuesBus,
72
73 idx_encoder: Arc<Encoder>,
75 min_cached_idx: usize,
76 max_cached: usize,
77 commit_mult: usize,
78
79 continuations_enabled: bool,
82}
83
84impl ProofShapeModule {
85 pub fn new(
86 mvk: &MultiStarkVkeyFrame,
87 b: &mut BusIndexManager,
88 bus_inventory: BusInventory,
89 continuations_enabled: bool,
90 ) -> Self {
91 assert!(
92 mvk.per_air.len() <= 1 << 8,
93 "recursion circuit only supports child verifying keys with at most 256 AIRs"
94 );
95
96 let idx_encoder = Arc::new(Encoder::new(mvk.per_air.len(), 2, true));
97
98 let (min_cached_idx, min_cached) = mvk
99 .per_air
100 .iter()
101 .enumerate()
102 .min_by_key(|(_, avk)| avk.params.width.cached_mains.len())
103 .map(|(idx, avk)| (idx, avk.params.width.cached_mains.len()))
104 .unwrap();
105 let mut max_cached = mvk
106 .per_air
107 .iter()
108 .map(|avk| avk.params.width.cached_mains.len())
109 .max()
110 .unwrap();
111 if min_cached == max_cached {
112 max_cached += 1;
113 }
114
115 let per_air = mvk
116 .per_air
117 .iter()
118 .map(|avk| AirMetadata {
119 is_required: avk.is_required,
120 need_rot: avk.params.need_rot,
121 num_public_values: avk.params.num_public_values,
122 num_interactions: avk.num_interactions,
123 main_width: avk.params.width.common_main,
124 cached_widths: avk.params.width.cached_mains.clone(),
125 preprocessed_width: avk.params.width.preprocessed,
126 preprocessed_data: avk.preprocessed_data.clone(),
127 })
128 .collect_vec();
129
130 let range_bus = bus_inventory.range_checker_bus;
131 let pow_bus = bus_inventory.power_checker_bus;
132 Self {
133 per_air,
134 l_skip: mvk.params.l_skip,
135 max_interaction_count: mvk.params.logup.max_interaction_count,
136 bus_inventory,
137 range_bus,
138 pow_bus,
139 permutation_bus: ProofShapePermutationBus::new(b.new_bus_idx()),
140 starting_tidx_bus: StartingTidxBus::new(b.new_bus_idx()),
141 num_pvs_bus: NumPublicValuesBus::new(b.new_bus_idx()),
142 idx_encoder,
143 min_cached_idx,
144 max_cached,
145 commit_mult: mvk.params.whir.rounds.first().unwrap().num_queries,
146 continuations_enabled,
147 }
148 }
149
150 #[tracing::instrument(level = "trace", skip_all)]
151 pub fn run_preflight<TS>(
152 &self,
153 child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
154 proof: &Proof<BabyBearPoseidon2Config>,
155 preflight: &mut Preflight,
156 ts: &mut TS,
157 ) where
158 TS: FiatShamirTranscript<BabyBearPoseidon2Config> + TranscriptHistory,
159 {
160 let l_skip = child_vk.inner.params.l_skip;
161 ts.observe_commit(child_vk.pre_hash);
162 ts.observe_commit(proof.common_main_commit);
163
164 let mut pvs_tidx = vec![];
165 let mut starting_tidx = vec![];
166
167 for (trace_vdata, avk, pvs) in izip!(
168 &proof.trace_vdata,
169 &child_vk.inner.per_air,
170 &proof.public_values
171 ) {
172 let is_air_present = trace_vdata.is_some();
173 starting_tidx.push(ts.len());
174
175 if !avk.is_required {
176 ts.observe(F::from_bool(is_air_present));
177 }
178 if let Some(trace_vdata) = trace_vdata {
179 if let Some(pdata) = avk.preprocessed_data.as_ref() {
180 ts.observe_commit(pdata.commit);
181 } else {
182 ts.observe(F::from_usize(trace_vdata.log_height));
183 }
184 debug_assert_eq!(avk.num_cached_mains(), trace_vdata.cached_commitments.len());
185 if !pvs.is_empty() {
186 pvs_tidx.push(ts.len());
187 }
188 for commit in &trace_vdata.cached_commitments {
189 ts.observe_commit(*commit);
190 }
191 debug_assert_eq!(avk.params.num_public_values, pvs.len());
192 }
193 for pv in pvs {
194 ts.observe(*pv);
195 }
196 }
197
198 let mut sorted_trace_vdata: Vec<_> = proof
199 .trace_vdata
200 .iter()
201 .cloned()
202 .enumerate()
203 .filter_map(|(air_id, data)| data.map(|data| (air_id, data)))
204 .collect();
205 sorted_trace_vdata.sort_by_key(|(air_idx, data)| (Reverse(data.log_height), *air_idx));
206
207 let n_max = proof
208 .trace_vdata
209 .iter()
210 .flat_map(|datum| {
211 datum
212 .as_ref()
213 .map(|datum| datum.log_height.saturating_sub(l_skip))
214 })
215 .max()
216 .unwrap();
217 let num_layers = proof.gkr_proof.claims_per_layer.len();
218 let n_logup = num_layers.saturating_sub(l_skip);
219
220 preflight.proof_shape = ProofShapePreflight {
221 sorted_trace_vdata,
222 starting_tidx,
223 pvs_tidx,
224 post_tidx: ts.len(),
225 n_max,
226 n_logup,
227 l_skip: child_vk.inner.params.l_skip,
228 };
229 }
230}
231
232impl AirModule for ProofShapeModule {
233 fn num_airs(&self) -> usize {
234 3
235 }
236
237 fn airs<SC: StarkProtocolConfig<F = F>>(&self) -> Vec<AirRef<SC>> {
238 let proof_shape_air = ProofShapeAir::<4, 8> {
239 per_air: self.per_air.clone(),
240 l_skip: self.l_skip,
241 min_cached_idx: self.min_cached_idx,
242 max_cached: self.max_cached,
243 commit_mult: self.commit_mult,
244 max_interaction_count: self.max_interaction_count,
245 idx_encoder: self.idx_encoder.clone(),
246 range_bus: self.range_bus,
247 pow_bus: self.pow_bus,
248 permutation_bus: self.permutation_bus,
249 starting_tidx_bus: self.starting_tidx_bus,
250 num_pvs_bus: self.num_pvs_bus,
251 fraction_folder_input_bus: self.bus_inventory.fraction_folder_input_bus,
252 expression_claim_n_max_bus: self.bus_inventory.expression_claim_n_max_bus,
253 gkr_module_bus: self.bus_inventory.gkr_module_bus,
254 air_shape_bus: self.bus_inventory.air_shape_bus,
255 air_presence_bus: self.bus_inventory.air_presence_bus,
256 hyperdim_bus: self.bus_inventory.hyperdim_bus,
257 lifted_heights_bus: self.bus_inventory.lifted_heights_bus,
258 commitments_bus: self.bus_inventory.commitments_bus,
259 transcript_bus: self.bus_inventory.transcript_bus,
260 n_lift_bus: self.bus_inventory.n_lift_bus,
261 eq_n_logup_n_max_bus: self.bus_inventory.eq_n_logup_n_max_bus,
262 eq_3b_shape_bus: self.bus_inventory.eq_3b_shape_bus,
263 cached_commit_bus: self.bus_inventory.cached_commit_bus,
264 pre_hash_bus: self.bus_inventory.pre_hash_bus,
265 continuations_enabled: self.continuations_enabled,
266 };
267 let pvs_air = PublicValuesAir {
268 public_values_bus: self.bus_inventory.public_values_bus,
269 num_pvs_bus: self.num_pvs_bus,
270 transcript_bus: self.bus_inventory.transcript_bus,
271 continuations_enabled: self.continuations_enabled,
272 };
273 let range_checker = RangeCheckerAir::<8> {
274 bus: self.range_bus,
275 };
276 vec![
277 Arc::new(proof_shape_air) as AirRef<_>,
278 Arc::new(pvs_air) as AirRef<_>,
279 Arc::new(range_checker) as AirRef<_>,
280 ]
281 }
282}
283
284impl<SC: StarkProtocolConfig<F = F>> TraceGenModule<GlobalCtxCpu, CpuBackend<SC>>
285 for ProofShapeModule
286{
287 type ModuleSpecificCtx<'a> = (
289 Arc<PowerCheckerCpuTraceGenerator<2, POW_CHECKER_HEIGHT>>,
290 &'a [usize],
291 );
292
293 #[tracing::instrument(skip_all)]
294 fn generate_proving_ctxs(
295 &self,
296 child_vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
297 proofs: &[Proof<BabyBearPoseidon2Config>],
298 preflights: &[Preflight],
299 ctx: &Self::ModuleSpecificCtx<'_>,
300 required_heights: Option<&[usize]>,
301 ) -> Option<Vec<AirProvingContext<CpuBackend<SC>>>> {
302 let pow_checker = &ctx.0;
303 let external_range_checks = ctx.1;
304
305 let range_checker = Arc::new(RangeCheckerCpuTraceGenerator::<8>::default());
306 let proof_shape = proof_shape::ProofShapeChip::<4, 8>::new(
307 self.idx_encoder.clone(),
308 self.min_cached_idx,
309 self.max_cached,
310 range_checker.clone(),
311 pow_checker.clone(),
312 );
313 let ctx = (child_vk, proofs, preflights);
314 let chips = [
315 ProofShapeModuleChip::ProofShape(proof_shape),
316 ProofShapeModuleChip::PublicValues,
317 ];
318 let mut ctxs: Vec<_> = chips
319 .par_iter()
320 .map(|chip| {
321 chip.generate_proving_ctx(
322 &ctx,
323 required_heights.map(|heights| heights[chip.index()]),
324 )
325 })
326 .collect::<Vec<_>>()
327 .into_iter()
328 .collect::<Option<Vec<_>>>()?;
329
330 for &val in external_range_checks {
331 range_checker.add_count(val);
332 }
333 tracing::trace_span!("wrapper.generate_trace", air = "RangeChecker").in_scope(|| {
334 ctxs.push(AirProvingContext::simple_no_pis(
335 range_checker.generate_trace_row_major(),
336 ));
337 });
338 Some(ctxs)
339 }
340}
341
342#[derive(strum_macros::Display, strum::EnumDiscriminants)]
343#[strum_discriminants(repr(usize))]
344enum ProofShapeModuleChip {
345 ProofShape(proof_shape::ProofShapeChip<4, 8>),
346 PublicValues,
347}
348
349impl ProofShapeModuleChip {
350 fn index(&self) -> usize {
351 ProofShapeModuleChipDiscriminants::from(self) as usize
352 }
353}
354
355impl RowMajorChip<F> for ProofShapeModuleChip {
356 type Ctx<'a> = (
357 &'a MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
358 &'a [Proof<BabyBearPoseidon2Config>],
359 &'a [Preflight],
360 );
361
362 #[tracing::instrument(
363 name = "wrapper.generate_trace",
364 level = "trace",
365 skip_all,
366 fields(air = %self)
367 )]
368 fn generate_trace(
369 &self,
370 ctx: &Self::Ctx<'_>,
371 required_height: Option<usize>,
372 ) -> Option<RowMajorMatrix<F>> {
373 use ProofShapeModuleChip::*;
374 match self {
375 ProofShape(chip) => chip.generate_trace(ctx, required_height),
376 PublicValues => {
377 pvs::PublicValuesTraceGenerator.generate_trace(&(ctx.1, ctx.2), required_height)
378 }
379 }
380 }
381}
382
383#[cfg(feature = "cuda")]
384mod cuda_tracegen {
385 use openvm_cuda_backend::GpuBackend;
386
387 use super::*;
388 use crate::{
389 cuda::{preflight::PreflightGpu, proof::ProofGpu, vk::VerifyingKeyGpu, GlobalCtxGpu},
390 primitives::{
391 pow::cuda::PowerCheckerGpuTraceGenerator, range::cuda::RangeCheckerGpuTraceGenerator,
392 },
393 };
394
395 impl TraceGenModule<GlobalCtxGpu, GpuBackend> for ProofShapeModule {
396 type ModuleSpecificCtx<'a> = (
397 Arc<PowerCheckerGpuTraceGenerator<2, POW_CHECKER_HEIGHT>>,
398 &'a [usize],
399 &'a openvm_cuda_common::stream::GpuDeviceCtx,
400 );
401
402 #[tracing::instrument(skip_all)]
403 fn generate_proving_ctxs(
404 &self,
405 child_vk: &VerifyingKeyGpu,
406 proofs: &[ProofGpu],
407 preflights: &[PreflightGpu],
408 ctx: &Self::ModuleSpecificCtx<'_>,
409 required_heights: Option<&[usize]>,
410 ) -> Option<Vec<AirProvingContext<GpuBackend>>> {
411 use crate::tracegen::ModuleChip;
412
413 let pow_checker_gpu = &ctx.0;
414 let external_range_checks = ctx.1;
415 let device_ctx = ctx.2;
416
417 let range_checker_gpu = Arc::new(RangeCheckerGpuTraceGenerator::<8>::from_vals(
418 external_range_checks,
419 device_ctx.clone(),
420 ));
421 let proof_shape_chip = proof_shape::cuda::ProofShapeChipGpu::<4, 8>::new(
422 self.idx_encoder.width(),
423 self.min_cached_idx,
424 self.max_cached,
425 range_checker_gpu.clone(),
426 pow_checker_gpu.clone(),
427 );
428 let mut ctxs = Vec::with_capacity(3);
429 let proof_shape_ctx =
434 tracing::trace_span!("wrapper.generate_trace", air = "ProofShape").in_scope(
435 || {
436 proof_shape_chip.generate_proving_ctx(
437 &(child_vk, preflights, device_ctx),
438 required_heights.map(|heights| heights[0]),
439 )
440 },
441 )?;
442 ctxs.push(proof_shape_ctx);
443
444 let public_values_ctx =
445 tracing::trace_span!("wrapper.generate_trace", air = "PublicValues").in_scope(
446 || {
447 pvs::cuda::PublicValuesGpuTraceGenerator.generate_proving_ctx(
448 &(proofs, preflights, device_ctx),
449 required_heights.map(|heights| heights[1]),
450 )
451 },
452 )?;
453 ctxs.push(public_values_ctx);
454 drop(proof_shape_chip);
457 tracing::trace_span!("wrapper.generate_trace", air = "RangeChecker").in_scope(|| {
460 ctxs.push(AirProvingContext::simple_no_pis(
461 Arc::try_unwrap(range_checker_gpu)
462 .ok()
463 .expect("range checker still shared")
464 .generate_trace(),
465 ));
466 });
467
468 Some(ctxs)
469 }
470 }
471}