Skip to main content

openvm_stark_backend/prover/
cpu_backend.rs

1//! CPU [ProverBackend] trait implementation.
2
3use 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    /// The random opening point `r` where the batch constraint sumcheck reduces to evaluation
108    /// claims of trace matrices `T, T_{rot}` at `r_{n_T}`.
109    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    /// The shared vector `r` where each trace matrix `T, T_{rot}` is opened at `r_{n_T}`.
135    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        // Currently alternates between preprocessed and cached pcs data
156        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}