1use getset::Getters;
5use itertools::Itertools;
6use openvm_stark_backend::{
7 keygen::types::MultiStarkProvingKey,
8 poly_common::Squarable,
9 proof::{BatchConstraintProof, GkrProof, StackingProof, WhirProof},
10 prover::{
11 poly::Mle,
12 stacked_pcs::{stacked_matrix, StackedPcsData},
13 stacked_reduction::prove_stacked_opening_reduction,
14 ColMajorMatrix, CommittedTraceData, DeviceDataTransporter, DeviceMultiStarkProvingKey,
15 DeviceStarkProvingKey, MatrixDimensions, MultiRapProver, OpeningProver, ProverDevice,
16 ProvingContext, StridedColMajorMatrixView, TraceCommitter,
17 },
18 FiatShamirTranscript, StarkProtocolConfig, SystemParams,
19};
20use p3_field::{ExtensionField, PrimeCharacteristicRing, TwoAdicField};
21use p3_matrix::dense::RowMajorMatrix;
22use p3_maybe_rayon::prelude::*;
23use tracing::instrument;
24
25use crate::{
26 backend::CpuBackend,
27 error::CpuProverError,
28 merkle::{rs_encode_and_merkle_cpu, CpuMerkleTree},
29 pcs_data::CpuStackedPcsData,
30 stacked_reduction::StackedReductionCpuNew,
31 two_adic::DftTwiddles,
32};
33
34#[derive(Clone, Getters, derive_new::new)]
36pub struct CpuDevice<SC> {
37 #[getset(get = "pub")]
38 config: SC,
39}
40
41impl<SC: StarkProtocolConfig> CpuDevice<SC> {
42 pub fn params(&self) -> &SystemParams {
43 self.config.params()
44 }
45}
46
47impl<SC, TS> ProverDevice<CpuBackend<SC>, TS> for CpuDevice<SC>
48where
49 SC: StarkProtocolConfig,
50 SC::F: Ord,
51 SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
52 TS: FiatShamirTranscript<SC>,
53{
54 type Error = CpuProverError;
55 type DeviceCtx = ();
56
57 fn device_ctx(&self) -> &() {
58 &()
59 }
60}
61
62impl<SC: StarkProtocolConfig> TraceCommitter<CpuBackend<SC>> for CpuDevice<SC>
63where
64 SC::F: Ord,
65{
66 type Error = CpuProverError;
67
68 #[instrument(level = "info", name = "trace_commit_cpu", skip_all)]
69 fn commit(
70 &self,
71 traces: &[&RowMajorMatrix<SC::F>],
72 ) -> Result<(SC::Digest, CpuStackedPcsData<SC::F, SC::Digest>), Self::Error> {
73 let col_major_traces: Vec<ColMajorMatrix<SC::F>> = traces
75 .iter()
76 .map(|rm| ColMajorMatrix::from_row_major(rm))
77 .collect();
78 let col_major_refs: Vec<&ColMajorMatrix<SC::F>> = col_major_traces.iter().collect();
79
80 let params = self.params();
81 let (q_trace, layout) = stacked_matrix(params.l_skip, params.n_stack, &col_major_refs)?;
82 let tree = rs_encode_and_merkle_cpu(
83 self.config().hasher(),
84 params.l_skip,
85 params.log_blowup,
86 &q_trace,
87 1 << params.k_whir(),
88 );
89 let root = tree.root()?;
90 let data = CpuStackedPcsData::new(layout, q_trace, tree);
91 Ok((root, data))
92 }
93}
94
95impl<SC, TS> MultiRapProver<CpuBackend<SC>, TS> for CpuDevice<SC>
96where
97 SC: StarkProtocolConfig,
98 SC::EF: TwoAdicField + ExtensionField<SC::F>,
99 TS: FiatShamirTranscript<SC>,
100{
101 type PartialProof = (GkrProof<SC>, BatchConstraintProof<SC>);
102 type Artifacts = Vec<SC::EF>;
103
104 type Error = CpuProverError;
105
106 fn prove_rap_constraints(
107 &self,
108 transcript: &mut TS,
109 mpk: &DeviceMultiStarkProvingKey<CpuBackend<SC>>,
110 ctx: &ProvingContext<CpuBackend<SC>>,
111 _common_main_pcs_data: &CpuStackedPcsData<SC::F, SC::Digest>,
112 ) -> Result<((GkrProof<SC>, BatchConstraintProof<SC>), Vec<SC::EF>), Self::Error> {
113 let (gkr_proof, batch_constraint_proof, r) =
114 crate::logup_zerocheck::prove_zerocheck_and_logup::<SC, _>(transcript, mpk, ctx)?;
115 Ok(((gkr_proof, batch_constraint_proof), r))
116 }
117}
118
119impl<SC, TS> OpeningProver<CpuBackend<SC>, TS> for CpuDevice<SC>
120where
121 SC: StarkProtocolConfig,
122 SC::F: Ord,
123 SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
124 TS: FiatShamirTranscript<SC>,
125{
126 type OpeningProof = (StackingProof<SC>, WhirProof<SC>);
127 type OpeningPoints = Vec<SC::EF>;
128
129 type Error = CpuProverError;
130
131 fn prove_openings(
132 &self,
133 transcript: &mut TS,
134 mpk: &DeviceMultiStarkProvingKey<CpuBackend<SC>>,
135 ctx: ProvingContext<CpuBackend<SC>>,
136 common_main_pcs_data: CpuStackedPcsData<SC::F, SC::Digest>,
137 r: Vec<SC::EF>,
138 ) -> Result<(StackingProof<SC>, WhirProof<SC>), Self::Error> {
139 let params = self.params();
140
141 let need_rot_per_trace = ctx
142 .per_trace
143 .iter()
144 .map(|(air_idx, _)| mpk.per_air[*air_idx].vk.params.need_rot)
145 .collect_vec();
146
147 let pre_cached_pcs_data_per_commit: Vec<_> = ctx
148 .per_trace
149 .iter()
150 .flat_map(|(air_idx, trace_ctx)| {
151 mpk.per_air[*air_idx]
152 .preprocessed_data
153 .iter()
154 .chain(&trace_ctx.cached_mains)
155 .map(|cd| cd.data.clone())
156 })
157 .collect();
158
159 let mut stacked_per_commit = vec![&common_main_pcs_data];
160 for data in &pre_cached_pcs_data_per_commit {
161 stacked_per_commit.push(data);
162 }
163 let mut need_rot_per_commit = vec![need_rot_per_trace];
164 for (air_idx, trace_ctx) in &ctx.per_trace {
165 let need_rot = mpk.per_air[*air_idx].vk.params.need_rot;
166 if mpk.per_air[*air_idx].preprocessed_data.is_some() {
167 need_rot_per_commit.push(vec![need_rot]);
168 }
169 for _ in &trace_ctx.cached_mains {
170 need_rot_per_commit.push(vec![need_rot]);
171 }
172 }
173 let (stacking_proof, u_prisma) =
174 prove_stacked_opening_reduction::<SC, _, _, _, StackedReductionCpuNew<SC>>(
175 self,
176 transcript,
177 params.n_stack,
178 stacked_per_commit,
179 need_rot_per_commit,
180 &r,
181 );
182
183 let (&u0, u_rest) = u_prisma
184 .split_first()
185 .ok_or(openvm_stark_backend::prover::error::WhirProverError::UPrismaEmpty)?;
186 let u_cube = u0
187 .exp_powers_of_2()
188 .take(params.l_skip)
189 .chain(u_rest.iter().copied())
190 .collect_vec();
191
192 let committed_mats = std::iter::once(&common_main_pcs_data)
194 .chain(pre_cached_pcs_data_per_commit.iter().map(|d| d.as_ref()))
195 .map(|d| (&d.matrix, &d.tree))
196 .collect_vec();
197
198 let whir_proof = crate::whir::prove_whir_opening_cpu::<SC, _>(
199 transcript,
200 self.config().hasher(),
201 params.l_skip,
202 params.log_blowup,
203 ¶ms.whir,
204 &committed_mats,
205 &u_cube,
206 )?;
207 Ok((stacking_proof, whir_proof))
208 }
209}
210
211impl<SC: StarkProtocolConfig> DeviceDataTransporter<SC, CpuBackend<SC>> for CpuDevice<SC> {
212 fn transport_pk_to_device(
213 &self,
214 mpk: &MultiStarkProvingKey<SC>,
215 ) -> DeviceMultiStarkProvingKey<CpuBackend<SC>> {
216 let per_air = mpk
217 .per_air
218 .iter()
219 .map(|pk| {
220 let preprocessed_data = pk.preprocessed_data.as_ref().map(|d| {
221 let view: StridedColMajorMatrixView<'_, SC::F> = d.mat_view(0).into();
222 let trace = view.to_row_major_matrix();
223 CommittedTraceData {
224 commitment: d.commit().unwrap(),
225 trace,
226 data: std::sync::Arc::new(stacked_pcs_data_to_cpu::<SC>(d)),
227 }
228 });
229 DeviceStarkProvingKey {
230 air_name: pk.air_name.clone(),
231 vk: pk.vk.clone(),
232 preprocessed_data,
233 other_data: (),
234 }
235 })
236 .collect();
237 DeviceMultiStarkProvingKey::new(
238 per_air,
239 mpk.trace_height_constraints.clone(),
240 mpk.max_constraint_degree,
241 mpk.params.clone(),
242 mpk.vk_pre_hash,
243 )
244 }
245
246 fn transport_matrix_to_device(&self, matrix: &ColMajorMatrix<SC::F>) -> RowMajorMatrix<SC::F> {
247 let view: StridedColMajorMatrixView<'_, SC::F> = matrix.as_view().into();
248 view.to_row_major_matrix()
249 }
250
251 fn transport_pcs_data_to_device(
252 &self,
253 pcs_data: &StackedPcsData<SC::F, SC::Digest>,
254 ) -> CpuStackedPcsData<SC::F, SC::Digest> {
255 stacked_pcs_data_to_cpu::<SC>(pcs_data)
256 }
257
258 fn transport_matrix_from_device_to_host(
259 &self,
260 matrix: &RowMajorMatrix<SC::F>,
261 ) -> ColMajorMatrix<SC::F> {
262 ColMajorMatrix::from_row_major(matrix)
263 }
264}
265
266fn stacked_pcs_data_to_cpu<SC: StarkProtocolConfig>(
270 pcs_data: &StackedPcsData<SC::F, SC::Digest>,
271) -> CpuStackedPcsData<SC::F, SC::Digest> {
272 let cm_backing = pcs_data.tree.backing_matrix();
273 let height = cm_backing.height();
274 let width = cm_backing.width();
275 let mut rm_values = SC::F::zero_vec(height * width);
277 rm_values
278 .par_chunks_exact_mut(width)
279 .enumerate()
280 .for_each(|(i, row)| {
281 for j in 0..width {
282 row[j] = cm_backing.values[j * height + i];
283 }
284 });
285 let rm_backing = RowMajorMatrix::new(rm_values, width);
286 let cpu_tree = unsafe {
287 CpuMerkleTree::from_raw_parts(
288 rm_backing,
289 pcs_data.tree.digest_layers().clone(),
290 pcs_data.tree.rows_per_query(),
291 )
292 };
293 CpuStackedPcsData::new(pcs_data.layout.clone(), pcs_data.matrix.clone(), cpu_tree)
294}
295
296pub(crate) fn eval_to_coeff_cpu<F: TwoAdicField>(evals: &[F], twiddles: &DftTwiddles<F>) -> Vec<F> {
300 let chunk_len = twiddles.size();
301 let mut buf = evals.to_vec();
302
303 for chunk in buf.chunks_exact_mut(chunk_len) {
305 twiddles.idft_inplace(chunk);
306 }
307
308 buf.par_chunks_exact_mut(chunk_len)
310 .for_each(Mle::coeffs_to_evals_inplace);
311
312 buf
313}