openvm_recursion_circuit/cuda/
proof.rs

1use itertools::Itertools;
2use openvm_cuda_common::{d_buffer::DeviceBuffer, stream::GpuDeviceCtx};
3use openvm_stark_backend::{keygen::types::MultiStarkVerifyingKey, proof::Proof};
4use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
5
6use crate::cuda::{to_device_or_nullptr_on, types::PublicValueData};
7
8/*
9 * Tracegen information (i.e. records) on a GPU device. Each field should
10 * be computable as soon as the verifier circuit has access to the child
11 * proof and verifying key.
12 */
13#[derive(Debug)]
14pub struct ProofGpu {
15    pub cpu: Proof<BabyBearPoseidon2Config>,
16    pub proof_shape: ProofShapeProofGpu,
17    pub gkr: GkrProofGpu,
18    pub batch_constraint: BatchConstraintProofGpu,
19    pub stacking: StackingProofGpu,
20    pub whir: WhirProofGpu,
21}
22
23#[derive(Debug)]
24pub struct ProofShapeProofGpu {
25    pub public_values: DeviceBuffer<PublicValueData>,
26}
27
28#[derive(Debug)]
29pub struct GkrProofGpu {
30    _dummy: usize,
31}
32
33#[derive(Debug)]
34pub struct BatchConstraintProofGpu {
35    _dummy: usize,
36}
37
38#[derive(Debug)]
39pub struct StackingProofGpu {
40    _dummy: usize,
41}
42
43#[derive(Debug)]
44pub struct WhirProofGpu {
45    _dummy: usize,
46}
47
48impl ProofGpu {
49    pub fn new(
50        vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
51        proof: &Proof<BabyBearPoseidon2Config>,
52        device_ctx: &GpuDeviceCtx,
53    ) -> Self {
54        ProofGpu {
55            cpu: proof.clone(),
56            proof_shape: Self::proof_shape(vk, proof, device_ctx),
57            gkr: Self::gkr(proof),
58            batch_constraint: Self::batch_constraint(proof),
59            stacking: Self::stacking(proof),
60            whir: Self::whir(proof),
61        }
62    }
63
64    fn proof_shape(
65        _vk: &MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
66        proof: &Proof<BabyBearPoseidon2Config>,
67        device_ctx: &GpuDeviceCtx,
68    ) -> ProofShapeProofGpu {
69        let num_airs = proof.public_values.len();
70        let public_values = proof
71            .public_values
72            .iter()
73            .enumerate()
74            .flat_map(move |(air_idx, pvs)| {
75                let air_num_pvs = pvs.len();
76                let total_airs = num_airs;
77                pvs.iter()
78                    .enumerate()
79                    .map(move |(pv_idx, &value)| PublicValueData {
80                        air_idx,
81                        air_num_pvs,
82                        num_airs: total_airs,
83                        pv_idx,
84                        value,
85                    })
86            })
87            .collect_vec();
88        ProofShapeProofGpu {
89            public_values: to_device_or_nullptr_on(&public_values, device_ctx).unwrap(),
90        }
91    }
92
93    fn gkr(_proof: &Proof<BabyBearPoseidon2Config>) -> GkrProofGpu {
94        GkrProofGpu { _dummy: 0 }
95    }
96
97    fn batch_constraint(_proof: &Proof<BabyBearPoseidon2Config>) -> BatchConstraintProofGpu {
98        BatchConstraintProofGpu { _dummy: 0 }
99    }
100
101    fn stacking(_proof: &Proof<BabyBearPoseidon2Config>) -> StackingProofGpu {
102        StackingProofGpu { _dummy: 0 }
103    }
104
105    fn whir(_proof: &Proof<BabyBearPoseidon2Config>) -> WhirProofGpu {
106        WhirProofGpu { _dummy: 0 }
107    }
108}