Skip to main content

openvm_cuda_backend/
gpu_backend.rs

1use std::marker::PhantomData;
2
3use itertools::Itertools;
4use openvm_cuda_common::memory_manager::MemTracker;
5use openvm_stark_backend::{
6    poly_common::Squarable,
7    proof::*,
8    prover::{
9        DeviceMultiStarkProvingKey, MultiRapProver, OpeningProver, ProverBackend, ProverDevice,
10        ProvingContext, TraceCommitter,
11    },
12};
13use tracing::instrument;
14
15use crate::{
16    base::DeviceMatrix,
17    hash_scheme::{DefaultHashScheme, GpuHashScheme},
18    logup_zerocheck::prove_zerocheck_and_logup_gpu,
19    merkle_tree::{MerkleProofQueryDigest, MerkleTreeConstructor},
20    prelude::{D_EF, EF, F},
21    sponge::GpuFiatShamirTranscript,
22    stacked_pcs::{stacked_commit, StackedPcsDataGpu},
23    stacked_reduction::prove_stacked_opening_reduction_gpu,
24    whir::prove_whir_opening_gpu,
25    AirDataGpu, GpuDevice, ProverError,
26};
27
28/// Generic GPU prover backend parameterised by a hash scheme `HS`.
29///
30/// Use the [`GpuBackend`] type alias to refer to the concrete BabyBear-Poseidon2
31/// backend without spelling out the generic parameter.
32#[derive(Clone, Copy)]
33pub struct GenericGpuBackend<HS: GpuHashScheme>(PhantomData<HS>);
34
35impl<HS: GpuHashScheme> Default for GenericGpuBackend<HS> {
36    fn default() -> Self {
37        Self(PhantomData)
38    }
39}
40
41/// Concrete GPU backend using the default BabyBear-Poseidon2 hash scheme.
42pub type GpuBackend = GenericGpuBackend<DefaultHashScheme>;
43
44impl<HS: GpuHashScheme> ProverBackend for GenericGpuBackend<HS> {
45    const CHALLENGE_EXT_DEGREE: u8 = D_EF as u8;
46
47    type Val = F;
48    type Challenge = EF;
49    type Commitment = HS::Digest;
50    type Matrix = DeviceMatrix<F>;
51    type PcsData = StackedPcsDataGpu<F, HS::Digest>;
52    type OtherAirData = AirDataGpu;
53}
54
55impl<HS: GpuHashScheme> TraceCommitter<GenericGpuBackend<HS>> for GpuDevice
56where
57    HS::MerkleHash: MerkleTreeConstructor,
58{
59    type Error = ProverError;
60
61    #[allow(clippy::type_complexity)]
62    #[instrument(name = "prover.commit", skip_all, fields(phase = "prover"))]
63    fn commit(
64        &self,
65        traces: &[&DeviceMatrix<F>],
66    ) -> Result<(HS::Digest, StackedPcsDataGpu<F, HS::Digest>), Self::Error> {
67        let cfg = self.params();
68        stacked_commit::<HS::MerkleHash>(
69            cfg.l_skip,
70            cfg.n_stack,
71            cfg.log_blowup,
72            cfg.k_whir(),
73            traces,
74            *self.prover_config(),
75            &self.device_ctx,
76        )
77    }
78}
79
80impl<HS: GpuHashScheme, TS: GpuFiatShamirTranscript<HS::SC>> ProverDevice<GenericGpuBackend<HS>, TS>
81    for GpuDevice
82where
83    HS::MerkleHash: MerkleTreeConstructor,
84    HS::Digest: MerkleProofQueryDigest,
85{
86    type Error = ProverError;
87    type DeviceCtx = openvm_cuda_common::stream::GpuDeviceCtx;
88
89    fn device_ctx(&self) -> &openvm_cuda_common::stream::GpuDeviceCtx {
90        &self.device_ctx
91    }
92}
93
94impl<HS: GpuHashScheme, TS: GpuFiatShamirTranscript<HS::SC>>
95    MultiRapProver<GenericGpuBackend<HS>, TS> for GpuDevice
96{
97    type PartialProof = (GkrProof<HS::SC>, BatchConstraintProof<HS::SC>);
98    /// The random opening point `r` where the batch constraint sumcheck reduces to evaluation
99    /// claims of trace matrices `T, T_{rot}` at `r_{n_T}`.
100    type Artifacts = Vec<EF>;
101    type Error = ProverError;
102
103    #[allow(clippy::type_complexity)]
104    #[instrument(name = "prover.rap_constraints", skip_all, fields(phase = "prover"))]
105    fn prove_rap_constraints(
106        &self,
107        transcript: &mut TS,
108        mpk: &DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
109        ctx: &ProvingContext<GenericGpuBackend<HS>>,
110        _common_main_pcs_data: &StackedPcsDataGpu<F, HS::Digest>,
111    ) -> Result<((GkrProof<HS::SC>, BatchConstraintProof<HS::SC>), Vec<EF>), Self::Error> {
112        let mem = MemTracker::start_and_reset_peak("prover.rap_constraints");
113        let save_memory = self.prover_config().zerocheck_save_memory;
114        // Threshold for monomial evaluation path based on proof type:
115        // - App proofs (log_blowup=1): higher threshold (512)
116        // - Recursion proofs: lower threshold (64)
117        let monomial_num_y_threshold = if self.params().log_blowup == 1 {
118            512
119        } else {
120            64
121        };
122        let (gkr_proof, batch_constraint_proof, r) = prove_zerocheck_and_logup_gpu::<HS, TS>(
123            transcript,
124            mpk,
125            ctx,
126            save_memory,
127            monomial_num_y_threshold,
128            self.sm_count(),
129            &self.device_ctx,
130        )?;
131        mem.emit_metrics();
132        Ok(((gkr_proof, batch_constraint_proof), r))
133    }
134}
135
136impl<HS: GpuHashScheme, TS: GpuFiatShamirTranscript<HS::SC>>
137    OpeningProver<GenericGpuBackend<HS>, TS> for GpuDevice
138where
139    HS::MerkleHash: MerkleTreeConstructor,
140    HS::Digest: MerkleProofQueryDigest,
141{
142    type OpeningProof = (StackingProof<HS::SC>, WhirProof<HS::SC>);
143    /// The shared vector `r` where each trace matrix `T, T_{rot}` is opened at `r_{n_T}`.
144    type OpeningPoints = Vec<EF>;
145    type Error = ProverError;
146
147    #[instrument(name = "prover.openings", skip_all, fields(phase = "prover"))]
148    fn prove_openings(
149        &self,
150        transcript: &mut TS,
151        mpk: &DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
152        ctx: ProvingContext<GenericGpuBackend<HS>>,
153        common_main_pcs_data: StackedPcsDataGpu<F, HS::Digest>,
154        r: Vec<EF>,
155    ) -> Result<Self::OpeningProof, Self::Error> {
156        let mut mem = MemTracker::start_and_reset_peak("prover.openings");
157        let params = self.params();
158        #[cfg(debug_assertions)]
159        {
160            let total_stacked_width: usize = std::iter::once(common_main_pcs_data.layout().width())
161                .chain(ctx.per_trace.iter().flat_map(|(air_idx, air_ctx)| {
162                    mpk.per_air[*air_idx]
163                        .preprocessed_data
164                        .iter()
165                        .map(|committed| committed.data.layout().width())
166                        .chain(
167                            air_ctx
168                                .cached_mains
169                                .iter()
170                                .map(|committed| committed.data.layout().width()),
171                        )
172                }))
173                .sum();
174            debug_assert!(
175                total_stacked_width <= mpk.params.w_stack,
176                "total stacked width across commits ({total_stacked_width}) exceeds w_stack ({})",
177                mpk.params.w_stack
178            );
179        }
180        let (stacking_proof, u_prisma, stacked_per_commit) =
181            prove_stacked_opening_reduction_gpu::<HS, TS>(
182                self,
183                transcript,
184                mpk,
185                ctx,
186                common_main_pcs_data,
187                &r,
188            )?;
189
190        let (&u0, u_rest) = u_prisma.split_first().unwrap();
191        let u_cube = u0
192            .exp_powers_of_2()
193            .take(params.l_skip)
194            .chain(u_rest.iter().copied())
195            .collect_vec();
196
197        let whir_proof = prove_whir_opening_gpu::<HS, TS>(
198            params,
199            transcript,
200            stacked_per_commit,
201            &u_cube,
202            &self.device_ctx,
203        )?;
204        mem.emit_metrics();
205        mem.reset_peak();
206        Ok((stacking_proof, whir_proof))
207    }
208}