use std::marker::PhantomData;
use getset::Getters;
use itertools::Itertools;
use p3_field::{ExtensionField, TwoAdicField};
use crate::{
keygen::types::MultiStarkProvingKey,
poly_common::Squarable,
proof::{BatchConstraintProof, GkrProof, StackingProof, WhirProof},
prover::{
error::RefProverError,
prove_zerocheck_and_logup,
stacked_pcs::{stacked_commit, StackedPcsData},
stacked_reduction::{prove_stacked_opening_reduction, StackedReductionCpu},
whir::WhirProver,
ColMajorMatrix, CommittedTraceData, DeviceDataTransporter, DeviceMultiStarkProvingKey,
DeviceStarkProvingKey, MultiRapProver, OpeningProver, ProverBackend, ProverDevice,
ProvingContext, TraceCommitter,
},
FiatShamirTranscript, StarkProtocolConfig, SystemParams,
};
#[derive(Clone, Copy)]
pub struct CpuColMajorBackend<SC: StarkProtocolConfig>(PhantomData<SC>);
impl<SC: StarkProtocolConfig> CpuColMajorBackend<SC> {
pub fn new() -> Self {
Self(PhantomData)
}
}
impl<SC: StarkProtocolConfig> Default for CpuColMajorBackend<SC> {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone, Getters, derive_new::new)]
pub struct ReferenceDevice<SC> {
#[getset(get = "pub")]
config: SC,
}
impl<SC: StarkProtocolConfig> ReferenceDevice<SC> {
pub fn params(&self) -> &SystemParams {
self.config.params()
}
}
impl<SC: StarkProtocolConfig> ProverBackend for CpuColMajorBackend<SC> {
const CHALLENGE_EXT_DEGREE: u8 = SC::D_EF as u8;
type Val = SC::F;
type Challenge = SC::EF;
type Commitment = SC::Digest;
type Matrix = ColMajorMatrix<SC::F>;
type OtherAirData = ();
type PcsData = StackedPcsData<SC::F, SC::Digest>;
}
impl<SC, TS> ProverDevice<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
where
SC: StarkProtocolConfig,
SC::F: Ord,
SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
TS: FiatShamirTranscript<SC>,
{
type Error = RefProverError;
type DeviceCtx = ();
fn device_ctx(&self) -> &() {
&()
}
}
impl<SC: StarkProtocolConfig> TraceCommitter<CpuColMajorBackend<SC>> for ReferenceDevice<SC>
where
SC::F: Ord,
{
type Error = RefProverError;
fn commit(
&self,
traces: &[&ColMajorMatrix<SC::F>],
) -> Result<(SC::Digest, StackedPcsData<SC::F, SC::Digest>), Self::Error> {
Ok(stacked_commit(
self.config().hasher(),
self.params().l_skip,
self.params().n_stack,
self.params().log_blowup,
self.params().k_whir(),
traces,
)?)
}
}
impl<SC, TS> MultiRapProver<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
where
SC: StarkProtocolConfig,
SC::EF: TwoAdicField + ExtensionField<SC::F>,
TS: FiatShamirTranscript<SC>,
{
type PartialProof = (GkrProof<SC>, BatchConstraintProof<SC>);
type Artifacts = Vec<SC::EF>;
type Error = RefProverError;
fn prove_rap_constraints(
&self,
transcript: &mut TS,
mpk: &DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
ctx: &ProvingContext<CpuColMajorBackend<SC>>,
_common_main_pcs_data: &StackedPcsData<SC::F, SC::Digest>,
) -> Result<((GkrProof<SC>, BatchConstraintProof<SC>), Vec<SC::EF>), Self::Error> {
let (gkr_proof, batch_constraint_proof, r) =
prove_zerocheck_and_logup::<SC, _>(transcript, mpk, ctx)?;
Ok(((gkr_proof, batch_constraint_proof), r))
}
}
impl<SC, TS> OpeningProver<CpuColMajorBackend<SC>, TS> for ReferenceDevice<SC>
where
SC: StarkProtocolConfig,
SC::F: Ord,
SC::EF: TwoAdicField + ExtensionField<SC::F> + Ord,
TS: FiatShamirTranscript<SC>,
{
type OpeningProof = (StackingProof<SC>, WhirProof<SC>);
type OpeningPoints = Vec<SC::EF>;
type Error = RefProverError;
fn prove_openings(
&self,
transcript: &mut TS,
mpk: &DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>>,
ctx: ProvingContext<CpuColMajorBackend<SC>>,
common_main_pcs_data: StackedPcsData<SC::F, SC::Digest>,
r: Vec<SC::EF>,
) -> Result<(StackingProof<SC>, WhirProof<SC>), Self::Error> {
let params = self.params();
let need_rot_per_trace = ctx
.per_trace
.iter()
.map(|(air_idx, _)| mpk.per_air[*air_idx].vk.params.need_rot)
.collect_vec();
let pre_cached_pcs_data_per_commit: Vec<_> = ctx
.per_trace
.iter()
.flat_map(|(air_idx, trace_ctx)| {
mpk.per_air[*air_idx]
.preprocessed_data
.iter()
.chain(&trace_ctx.cached_mains)
.map(|cd| cd.data.clone())
})
.collect();
let mut stacked_per_commit = vec![&common_main_pcs_data];
for data in &pre_cached_pcs_data_per_commit {
stacked_per_commit.push(data);
}
#[cfg(debug_assertions)]
{
let total_stacked_width: usize =
stacked_per_commit.iter().map(|d| d.layout.width()).sum();
debug_assert!(
total_stacked_width <= params.w_stack,
"total stacked width across commits ({total_stacked_width}) exceeds w_stack ({})",
params.w_stack
);
}
let mut need_rot_per_commit = vec![need_rot_per_trace];
for (air_idx, trace_ctx) in &ctx.per_trace {
let need_rot = mpk.per_air[*air_idx].vk.params.need_rot;
if mpk.per_air[*air_idx].preprocessed_data.is_some() {
need_rot_per_commit.push(vec![need_rot]);
}
for _ in &trace_ctx.cached_mains {
need_rot_per_commit.push(vec![need_rot]);
}
}
let (stacking_proof, u_prisma) =
prove_stacked_opening_reduction::<SC, _, _, _, StackedReductionCpu<SC>>(
self,
transcript,
params.n_stack,
stacked_per_commit,
need_rot_per_commit,
&r,
);
let (&u0, u_rest) = u_prisma
.split_first()
.ok_or(crate::prover::error::WhirProverError::UPrismaEmpty)?;
let u_cube = u0
.exp_powers_of_2()
.take(params.l_skip)
.chain(u_rest.iter().copied())
.collect_vec();
let whir_proof = self.prove_whir(
transcript,
common_main_pcs_data,
pre_cached_pcs_data_per_commit,
&u_cube,
)?;
Ok((stacking_proof, whir_proof))
}
}
impl<SC: StarkProtocolConfig> DeviceDataTransporter<SC, CpuColMajorBackend<SC>>
for ReferenceDevice<SC>
{
fn transport_pk_to_device(
&self,
mpk: &MultiStarkProvingKey<SC>,
) -> DeviceMultiStarkProvingKey<CpuColMajorBackend<SC>> {
let per_air = mpk
.per_air
.iter()
.map(|pk| {
let preprocessed_data = pk.preprocessed_data.as_ref().map(|d| {
let trace = d.mat_view(0).to_matrix();
CommittedTraceData {
commitment: d.commit().unwrap(),
trace,
data: d.clone(),
}
});
DeviceStarkProvingKey {
air_name: pk.air_name.clone(),
vk: pk.vk.clone(),
preprocessed_data,
other_data: (),
}
})
.collect();
DeviceMultiStarkProvingKey::new(
per_air,
mpk.trace_height_constraints.clone(),
mpk.max_constraint_degree,
mpk.params.clone(),
mpk.vk_pre_hash,
)
}
fn transport_matrix_to_device(&self, matrix: &ColMajorMatrix<SC::F>) -> ColMajorMatrix<SC::F> {
matrix.clone()
}
fn transport_pcs_data_to_device(
&self,
pcs_data: &StackedPcsData<SC::F, SC::Digest>,
) -> StackedPcsData<SC::F, SC::Digest> {
pcs_data.clone()
}
fn transport_matrix_from_device_to_host(
&self,
matrix: &ColMajorMatrix<SC::F>,
) -> ColMajorMatrix<SC::F> {
matrix.clone()
}
}