openvm_stark_backend/prover/
cpu_backend.rs1use std::marker::PhantomData;
4
5use getset::Getters;
6use itertools::Itertools;
7use p3_field::{ExtensionField, TwoAdicField};
8
9use crate::{
10 keygen::types::MultiStarkProvingKey,
11 poly_common::Squarable,
12 proof::{BatchConstraintProof, GkrProof, StackingProof, WhirProof},
13 prover::{
14 error::RefProverError,
15 prove_zerocheck_and_logup,
16 stacked_pcs::{stacked_commit, StackedPcsData},
17 stacked_reduction::{prove_stacked_opening_reduction, StackedReductionCpu},
18 whir::WhirProver,
19 ColMajorMatrix, CommittedTraceData, DeviceDataTransporter, DeviceMultiStarkProvingKey,
20 DeviceStarkProvingKey, MultiRapProver, OpeningProver, ProverBackend, ProverDevice,
21 ProvingContext, TraceCommitter,
22 },
23 FiatShamirTranscript, StarkProtocolConfig, SystemParams,
24};
25
26#[derive(Clone, Copy)]
27pub struct CpuColMajorBackend<SC: StarkProtocolConfig>(PhantomData<SC>);
28
29impl<SC: StarkProtocolConfig> CpuColMajorBackend<SC> {
30 pub fn new() -> Self {
31 Self(PhantomData)
32 }
33}
34
35impl<SC: StarkProtocolConfig> Default for CpuColMajorBackend<SC> {
36 fn default() -> Self {
37 Self::new()
38 }
39}
40
41#[derive(Clone, Getters, derive_new::new)]
42pub struct ReferenceDevice<SC> {
43 #[getset(get = "pub")]
44 config: SC,
45}
46
47impl<SC: StarkProtocolConfig> ReferenceDevice<SC> {
48 pub fn params(&self) -> &SystemParams {
49 self.config.params()
50 }
51}
52
53impl<SC: StarkProtocolConfig> ProverBackend for CpuColMajorBackend<SC> {
54 const CHALLENGE_EXT_DEGREE: u8 = SC::D_EF as u8;
55
56 type Val = SC::F;
57 type Challenge = SC::EF;
58 type Commitment = SC::Digest;
59 type Matrix = ColMajorMatrix<SC::F>;
60 type OtherAirData = ();
61 type PcsData = StackedPcsData<SC::F, SC::Digest>;
62}
63
64impl<SC, TS> ProverDevice<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
65where
66 SC: StarkProtocolConfig,
67 SC::F: Ord,
68 SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
69 TS: FiatShamirTranscript<SC>,
70{
71 type Error = RefProverError;
72 type DeviceCtx = ();
73
74 fn device_ctx(&self) -> &() {
75 &()
76 }
77}
78
79impl<SC: StarkProtocolConfig> TraceCommitter<CpuColMajorBackend<SC>> for ReferenceDevice<SC>
80where
81 SC::F: Ord,
82{
83 type Error = RefProverError;
84
85 fn commit(
86 &self,
87 traces: &[&ColMajorMatrix<SC::F>],
88 ) -> Result<(SC::Digest, StackedPcsData<SC::F, SC::Digest>), Self::Error> {
89 Ok(stacked_commit(
90 self.config().hasher(),
91 self.params().l_skip,
92 self.params().n_stack,
93 self.params().log_blowup,
94 self.params().k_whir(),
95 traces,
96 )?)
97 }
98}
99
100impl<SC, TS> MultiRapProver<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
101where
102 SC: StarkProtocolConfig,
103 SC::EF: TwoAdicField + ExtensionField<SC::F>,
104 TS: FiatShamirTranscript<SC>,
105{
106 type PartialProof = (GkrProof<SC>, BatchConstraintProof<SC>);
107 type Artifacts = Vec<SC::EF>;
110
111 type Error = RefProverError;
112
113 fn prove_rap_constraints(
114 &self,
115 transcript: &mut TS,
116 mpk: &DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
117 ctx: &ProvingContext<CpuColMajorBackend<SC>>,
118 _common_main_pcs_data: &StackedPcsData<SC::F, SC::Digest>,
119 ) -> Result<((GkrProof<SC>, BatchConstraintProof<SC>), Vec<SC::EF>), Self::Error> {
120 let (gkr_proof, batch_constraint_proof, r) =
121 prove_zerocheck_and_logup::<SC, _>(transcript, mpk, ctx)?;
122 Ok(((gkr_proof, batch_constraint_proof), r))
123 }
124}
125
126impl<SC, TS> OpeningProver<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
127where
128 SC: StarkProtocolConfig,
129 SC::F: Ord,
130 SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
131 TS: FiatShamirTranscript<SC>,
132{
133 type OpeningProof = (StackingProof<SC>, WhirProof<SC>);
134 type OpeningPoints = Vec<SC::EF>;
136
137 type Error = RefProverError;
138
139 fn prove_openings(
140 &self,
141 transcript: &mut TS,
142 mpk: &DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
143 ctx: ProvingContext<CpuColMajorBackend<SC>>,
144 common_main_pcs_data: StackedPcsData<SC::F, SC::Digest>,
145 r: Vec<SC::EF>,
146 ) -> Result<(StackingProof<SC>, WhirProof<SC>), Self::Error> {
147 let params = self.params();
148
149 let need_rot_per_trace = ctx
150 .per_trace
151 .iter()
152 .map(|(air_idx, _)| mpk.per_air[*air_idx].vk.params.need_rot)
153 .collect_vec();
154
155 let pre_cached_pcs_data_per_commit: Vec<_> = ctx
157 .per_trace
158 .iter()
159 .flat_map(|(air_idx, trace_ctx)| {
160 mpk.per_air[*air_idx]
161 .preprocessed_data
162 .iter()
163 .chain(&trace_ctx.cached_mains)
164 .map(|cd| cd.data.clone())
165 })
166 .collect();
167
168 let mut stacked_per_commit = vec![&common_main_pcs_data];
169 for data in &pre_cached_pcs_data_per_commit {
170 stacked_per_commit.push(data);
171 }
172 #[cfg(debug_assertions)]
173 {
174 let total_stacked_width: usize =
175 stacked_per_commit.iter().map(|d| d.layout.width()).sum();
176 debug_assert!(
177 total_stacked_width <= params.w_stack,
178 "total stacked width across commits ({total_stacked_width}) exceeds w_stack ({})",
179 params.w_stack
180 );
181 }
182
183 let mut need_rot_per_commit = vec![need_rot_per_trace];
184 for (air_idx, trace_ctx) in &ctx.per_trace {
185 let need_rot = mpk.per_air[*air_idx].vk.params.need_rot;
186 if mpk.per_air[*air_idx].preprocessed_data.is_some() {
187 need_rot_per_commit.push(vec![need_rot]);
188 }
189 for _ in &trace_ctx.cached_mains {
190 need_rot_per_commit.push(vec![need_rot]);
191 }
192 }
193 let (stacking_proof, u_prisma) =
194 prove_stacked_opening_reduction::<SC, _, _, _, StackedReductionCpu<SC>>(
195 self,
196 transcript,
197 params.n_stack,
198 stacked_per_commit,
199 need_rot_per_commit,
200 &r,
201 );
202
203 let (&u0, u_rest) = u_prisma
204 .split_first()
205 .ok_or(crate::prover::error::WhirProverError::UPrismaEmpty)?;
206 let u_cube = u0
207 .exp_powers_of_2()
208 .take(params.l_skip)
209 .chain(u_rest.iter().copied())
210 .collect_vec();
211
212 let whir_proof = self.prove_whir(
213 transcript,
214 common_main_pcs_data,
215 pre_cached_pcs_data_per_commit,
216 &u_cube,
217 )?;
218 Ok((stacking_proof, whir_proof))
219 }
220}
221
222impl<SC: StarkProtocolConfig> DeviceDataTransporter<SC, CpuColMajorBackend<SC>>
223 for ReferenceDevice<SC>
224{
225 fn transport_pk_to_device(
226 &self,
227 mpk: &MultiStarkProvingKey<SC>,
228 ) -> DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>> {
229 let per_air = mpk
230 .per_air
231 .iter()
232 .map(|pk| {
233 let preprocessed_data = pk.preprocessed_data.as_ref().map(|d| {
234 let trace = d.mat_view(0).to_matrix();
235 CommittedTraceData {
236 commitment: d.commit().unwrap(),
237 trace,
238 data: d.clone(),
239 }
240 });
241 DeviceStarkProvingKey {
242 air_name: pk.air_name.clone(),
243 vk: pk.vk.clone(),
244 preprocessed_data,
245 other_data: (),
246 }
247 })
248 .collect();
249 DeviceMultiStarkProvingKey::new(
250 per_air,
251 mpk.trace_height_constraints.clone(),
252 mpk.max_constraint_degree,
253 mpk.params.clone(),
254 mpk.vk_pre_hash,
255 )
256 }
257
258 fn transport_matrix_to_device(&self, matrix: &ColMajorMatrix<SC::F>) -> ColMajorMatrix<SC::F> {
259 matrix.clone()
260 }
261
262 fn transport_pcs_data_to_device(
263 &self,
264 pcs_data: &StackedPcsData<SC::F, SC::Digest>,
265 ) -> StackedPcsData<SC::F, SC::Digest> {
266 pcs_data.clone()
267 }
268
269 fn transport_matrix_from_device_to_host(
270 &self,
271 matrix: &ColMajorMatrix<SC::F>,
272 ) -> ColMajorMatrix<SC::F> {
273 matrix.clone()
274 }
275}