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#[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
41pub 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 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 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 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}