openvm_recursion_circuit/proof_shape/
mod.rs

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    // Verifying key fields
58    per_air: Vec<AirMetadata>,
59    l_skip: usize,
60    /// Threshold from the child VK used by [`ProofShapeAir`] on the summary row:
61    /// `sum_i(num_interactions[i] * lifted_height[i]) < max_interaction_count`,
62    /// with `lifted_height[i] = max(trace_height[i], 2^l_skip)`.
63    max_interaction_count: u32,
64
65    // Buses (inventory for external, others are internal)
66    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    // Required for ProofShapeAir tracegen + constraints
74    idx_encoder: Arc<Encoder>,
75    min_cached_idx: usize,
76    max_cached: usize,
77    commit_mult: usize,
78
79    // Module sends extra public values message for use outside of verifier
80    // sub-circuit if true
81    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    // (pow_checker, external_range_checks)
288    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            // PERF[jpw]: we avoid par_iter so that kernel launches occur on the same stream.
430            // This can be parallelized to separate streams for more CUDA stream parallelism, but it
431            // will require recording events so streams properly sync for cudaMemcpyAsync and kernel
432            // launches
433            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 the proof_shape chip so we can finalize auxiliary trace state (it holds Arc
455            // clones).
456            drop(proof_shape_chip);
457            // Caution: proof_shape **must** finish trace gen before we materialize range checker
458            // trace or sync power checker multiplicities to CPU.
459            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}