Skip to main content

openvm_cpu_backend/
device.rs

1//! [CpuDevice] implementation: TraceCommitter, DeviceDataTransporter, MultiRapProver,
2//! OpeningProver.
3
4use 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/// Row-major CPU prover device.
35#[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        // Convert row-major to col-major for stacking commitment
74        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        // Convert to col-major for WHIR
193        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            &params.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
266/// Convert stark-backend's `StackedPcsData` (ColMajor backing) to cpu-backend's
267/// `CpuStackedPcsData` (RowMajor backing). The eval matrix is cloned as-is (ColMajor),
268/// while the Merkle tree's backing matrix is transposed to RowMajor.
269fn 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    // Transpose ColMajor backing → RowMajor backing
276    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
296/// In-place PLE evaluation-to-coefficient conversion.
297/// Uses inline DIF iDFT with reusable twiddle factors to eliminate the per-chunk allocation
298/// overhead of the shared `eval_to_coeff_rs_message`.
299pub(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    // Phase 1: In-place DIF iDFT on each chunk.
304    for chunk in buf.chunks_exact_mut(chunk_len) {
305        twiddles.idft_inplace(chunk);
306    }
307
308    // Phase 2: Convert MLE coefficients to evaluations in-place.
309    buf.par_chunks_exact_mut(chunk_len)
310        .for_each(Mle::coeffs_to_evals_inplace);
311
312    buf
313}