use std::array::from_fn;
use cfg_if::cfg_if;
use itertools::Itertools;
use p3_dft::TwoAdicSubgroupDft;
use p3_field::{
batch_multiplicative_inverse, ExtensionField, Field, PrimeCharacteristicRing, TwoAdicField,
};
use p3_interpolation::interpolate_coset_with_precomputation;
use p3_matrix::dense::RowMajorMatrix;
use p3_maybe_rayon::prelude::*;
use p3_util::log2_strict_usize;
use tracing::{debug, instrument, trace};
use crate::{
dft::Radix2BowersSerial,
poly_common::UnivariatePoly,
prover::{
error::SumcheckError, ColMajorMatrix, ColMajorMatrixView, MatrixDimensions, MatrixView,
StridedColMajorMatrixView,
},
FiatShamirTranscript, StarkProtocolConfig,
};
#[instrument(level = "trace", skip_all)]
pub fn sumcheck_uni_round0_poly<F, EF, FN, const WD: usize>(
l_skip: usize,
n: usize,
d: usize,
mats: &[(StridedColMajorMatrixView<F>, bool)],
w: FN,
) -> [UnivariatePoly<EF>; WD]
where
F: TwoAdicField,
EF: ExtensionField<F> + TwoAdicField,
FN: Fn(
F,
usize,
&[Vec<F>],
) -> [EF; WD]
+ Sync,
{
if d == 0 {
return from_fn(|_| UnivariatePoly(vec![]));
}
#[cfg(debug_assertions)]
if n > 0 {
for (m, _) in mats.iter() {
assert_eq!(m.height(), 1 << (l_skip + n));
}
} else {
for (m, _) in mats.iter() {
assert!(
m.height() <= 1 << l_skip,
"mat height {} > 2^{l_skip}",
m.height()
);
}
}
let g = F::GENERATOR;
let omega_skip = F::two_adic_generator(l_skip);
let coset_shifts = g.powers().skip(1).take(d).collect_vec();
let evals = (0..1 << n).into_par_iter().map(|x| {
let dft = Radix2BowersSerial;
let mats_at_zs = mats
.iter()
.map(|(mat, is_rot)| {
let height = mat.height();
let offset = usize::from(*is_rot);
(0..mat.width())
.map(|col_idx| {
let col_x = ((x << l_skip)..(x + 1) << l_skip)
.map(|i| unsafe { *mat.get_unchecked((i + offset) % height, col_idx) })
.collect_vec();
let coeffs = dft.idft(col_x);
coset_shifts
.iter()
.flat_map(|&shift| dft.coset_dft(coeffs.clone(), shift))
.collect_vec()
})
.collect_vec()
})
.collect_vec();
omega_skip
.powers()
.take(1 << l_skip)
.enumerate()
.flat_map(|(z_idx, z)| {
coset_shifts
.iter()
.enumerate()
.map(|(coset_idx, &shift)| {
let z_int = (coset_idx << l_skip) + z_idx;
let row_z_x = mats_at_zs
.iter()
.map(|mat_at_zs| {
mat_at_zs
.iter()
.map(|col_at_zs| col_at_zs[z_int])
.collect_vec()
})
.collect_vec();
w(shift * z, x, &row_z_x)
})
.collect_vec()
})
.collect_vec()
});
let hypercube_sum = |mut acc: Vec<[EF; WD]>, x| {
for (acc, x) in acc.iter_mut().zip(x) {
for (acc_i, x_i) in acc.iter_mut().zip(x) {
*acc_i += x_i;
}
}
acc
};
cfg_if! {
if #[cfg(feature = "parallel")] {
let evals = evals.reduce(
|| vec![[EF::ZERO; WD]; d << l_skip],
hypercube_sum
);
} else {
let evals = evals.collect_vec();
let evals = evals.into_iter().fold(
vec![[EF::ZERO; WD]; d << l_skip],
hypercube_sum
);
}
}
from_fn(|i| {
let values = evals.iter().map(|x| x[i]).collect_vec();
UnivariatePoly::from_geometric_cosets_evals_idft(RowMajorMatrix::new(values, d), g, g)
})
}
pub const fn sumcheck_round0_deg(l_skip: usize, d: usize) -> usize {
d * ((1 << l_skip) - 1)
}
#[instrument(level = "trace", skip_all)]
pub fn fold_ple_evals<F, EF>(
l_skip: usize,
mat: StridedColMajorMatrixView<F>,
is_rot: bool,
r: EF,
) -> ColMajorMatrix<EF>
where
F: TwoAdicField,
EF: ExtensionField<F> + TwoAdicField,
{
let height = mat.height();
let lifted_height = height.max(1 << l_skip);
let width = mat.width();
let omega = F::two_adic_generator(l_skip);
let omega_pows = omega.powers().take(1 << l_skip).collect_vec();
let denoms = omega_pows
.iter()
.map(|&x_i| r - EF::from(x_i))
.collect_vec();
let inv_denoms = batch_multiplicative_inverse(&denoms);
let offset = usize::from(is_rot);
let new_height = lifted_height >> l_skip;
let values = (0..width * new_height)
.into_par_iter()
.map(|idx| {
let x = idx % new_height;
let j = idx / new_height;
let uni_evals = (0..1 << l_skip)
.map(|z| unsafe { *mat.get_unchecked(((x << l_skip) + z + offset) % height, j) })
.collect_vec();
interpolate_coset_with_precomputation(
&RowMajorMatrix::new_col(uni_evals),
F::ONE,
r,
&omega_pows,
&inv_denoms,
)[0]
})
.collect::<Vec<_>>();
ColMajorMatrix::new(values, width)
}
pub fn batch_fold_ple_evals<F, EF>(
l_skip: usize,
mats: Vec<ColMajorMatrix<F>>,
is_rot: bool,
r: EF,
) -> Vec<ColMajorMatrix<EF>>
where
F: TwoAdicField,
EF: ExtensionField<F> + TwoAdicField,
{
mats.into_par_iter()
.map(|mat| fold_ple_evals(l_skip, mat.as_view().into(), is_rot, r))
.collect()
}
#[instrument(level = "trace", skip_all)]
pub fn sumcheck_round_poly_evals<F, FN, const WD: usize>(
n: usize,
d: usize,
mats: &[ColMajorMatrixView<F>],
w: FN,
) -> [Vec<F>; WD]
where
F: Field,
FN: Fn(
F,
usize,
&[Vec<F>],
) -> [F; WD]
+ Sync,
{
debug_assert!(mats.iter().all(|mat| mat.height() == 1 << n));
if n == 0 {
let evals = mats.iter().map(|row| row.values.to_vec()).collect_vec();
return w(F::ONE, 0, &evals).map(|x| vec![x; d]);
}
let hypercube_dim = n - 1;
let f_hat = |x: usize, y: usize| {
let x = F::from_usize(x);
let row_x_y = mats
.iter()
.map(|mat| {
mat.columns()
.map(|col| {
let t_0 = col[y << 1];
let t_1 = col[(y << 1) | 1];
t_0 + (t_1 - t_0) * x
})
.collect_vec()
})
.collect_vec();
w(x, y, &row_x_y)
};
trace!(sum_claim = ?{(0..1 << n)
.map(|x| f_hat(x & 1, x >> 1))
.fold([F::ZERO; WD], |mut acc, x| {
for (acc_i, x_i) in acc.iter_mut().zip(x) {
*acc_i += x_i;
}
acc
})
}, "sumcheck_round");
let evals = (0..1 << hypercube_dim)
.into_par_iter()
.map(|y| (1..=d).map(|x| f_hat(x, y)).collect_vec());
let hypercube_sum = |mut acc: Vec<[F; WD]>, x| {
for (acc, x) in acc.iter_mut().zip(x) {
for (acc_i, x_i) in acc.iter_mut().zip(x) {
*acc_i += x_i;
}
}
acc
};
cfg_if! {
if #[cfg(feature = "parallel")] {
let evals = evals.reduce(
|| vec![[F::ZERO; WD]; d],
hypercube_sum
);
} else {
let evals = evals.collect_vec();
let evals = evals.into_iter().fold(
vec![[F::ZERO; WD]; d],
hypercube_sum
);
}
}
from_fn(|i| evals.iter().map(|eval| eval[i]).collect_vec())
}
#[instrument(level = "trace", skip_all)]
pub fn fold_mle_evals<EF: Field>(mat: ColMajorMatrix<EF>, r: EF) -> ColMajorMatrix<EF> {
let height = mat.height();
if height <= 1 {
return mat;
}
let width = mat.width();
let values = mat
.values
.par_chunks_exact(height)
.flat_map(|t| {
t.par_chunks_exact(2).map(|t_01| {
let t_0 = t_01[0];
let t_1 = t_01[1];
t_0 + (t_1 - t_0) * r
})
})
.collect::<Vec<_>>();
ColMajorMatrix::new(values, width)
}
pub fn batch_fold_mle_evals<EF: Field>(
mats: Vec<ColMajorMatrix<EF>>,
r: EF,
) -> Vec<ColMajorMatrix<EF>> {
mats.into_par_iter()
.map(|mat| fold_mle_evals(mat, r))
.collect()
}
pub fn fold_mle_evals_inplace<EF: Field>(mat: &mut ColMajorMatrix<EF>, r: EF) {
let height = mat.height();
if height <= 1 {
return;
}
mat.values.par_chunks_exact_mut(height).for_each(|t| {
for y in 0..height / 2 {
let t_0 = t[y << 1];
let t_1 = t[(y << 1) + 1];
t[y] = t_0 + (t_1 - t_0) * r;
}
});
}
pub struct SumcheckCubeProof<EF> {
pub sum_claim: EF,
pub round_polys_eval: Vec<Vec<EF>>,
pub eval_claim: EF,
}
pub struct SumcheckPrismProof<EF> {
pub sum_claim: EF,
pub s_0: UnivariatePoly<EF>,
pub round_polys_eval: Vec<Vec<EF>>,
pub eval_claim: EF,
}
#[allow(clippy::type_complexity)]
pub fn sumcheck_multilinear<SC: StarkProtocolConfig, F: Field, TS: FiatShamirTranscript<SC>>(
transcript: &mut TS,
evals: &[F],
) -> Result<(SumcheckCubeProof<SC::EF>, Vec<SC::EF>), SumcheckError>
where
SC::EF: ExtensionField<F>,
{
let n = log2_strict_usize(evals.len());
let mut round_polys_eval = Vec::with_capacity(n);
let mut r = Vec::with_capacity(n);
let mut current_evals =
ColMajorMatrix::new(evals.iter().map(|&x| SC::EF::from(x)).collect(), 1);
let sum_claim: SC::EF = evals.iter().fold(F::ZERO, |acc, &x| acc + x).into();
transcript.observe_ext(sum_claim);
for round in 0..n {
let [s] =
sumcheck_round_poly_evals(n - round, 1, &[current_evals.as_view()], |_x, _y, evals| {
[evals[0][0]]
});
if s.len() != 1 {
return Err(SumcheckError::MultilinearRoundPolyLen { len: s.len() });
}
transcript.observe_ext(s[0]);
round_polys_eval.push(s);
let r_round = transcript.sample_ext();
debug!(%round, %r_round);
r.push(r_round);
current_evals = fold_mle_evals(current_evals, r_round);
}
if current_evals.values.len() != 1 {
return Err(SumcheckError::MultilinearFinalEvalLen {
len: current_evals.values.len(),
});
}
let eval_claim = current_evals.values[0];
transcript.observe_ext(eval_claim);
Ok((
SumcheckCubeProof {
sum_claim,
round_polys_eval,
eval_claim,
},
r,
))
}
#[allow(clippy::type_complexity)]
pub fn sumcheck_prismalinear<SC: StarkProtocolConfig, F, TS: FiatShamirTranscript<SC>>(
transcript: &mut TS,
l_skip: usize,
evals: &[F],
) -> Result<(SumcheckPrismProof<SC::EF>, Vec<SC::EF>), SumcheckError>
where
F: TwoAdicField,
SC::EF: ExtensionField<F> + TwoAdicField,
{
let prism_dim = log2_strict_usize(evals.len());
if prism_dim < l_skip {
return Err(SumcheckError::PrismalinearDimTooSmall { prism_dim, l_skip });
}
let n = prism_dim - l_skip;
let mut round_polys_eval = Vec::with_capacity(n);
let mut r = Vec::with_capacity(n + 1);
let sum_claim: SC::EF = evals.iter().copied().sum::<F>().into();
transcript.observe_ext(sum_claim);
let current_evals = ColMajorMatrix::new(evals.to_vec(), 1);
let [s_0] = sumcheck_uni_round0_poly(
l_skip,
n,
1,
&[(current_evals.as_view().into(), false)],
|_z, _x, evals| [evals[0][0]],
);
let s_0_ext = UnivariatePoly::new(
s_0.0
.into_iter()
.map(|x| {
let ext = SC::EF::from(x);
transcript.observe_ext(ext);
ext
})
.collect(),
);
let r_0 = transcript.sample_ext();
debug!(round = 0, r_round = %r_0);
r.push(r_0);
let mut current_evals = fold_ple_evals(l_skip, current_evals.as_view().into(), false, r_0);
debug_assert_eq!(current_evals.height(), 1 << n);
for round in 1..=n {
debug!(
cur_sum = %current_evals
.values
.iter()
.fold(SC::EF::ZERO, |acc, x| acc + *x)
);
let [s] = sumcheck_round_poly_evals(
n + 1 - round,
1,
&[current_evals.as_view()],
|_x, _y, evals| [evals[0][0]],
);
if s.len() != 1 {
return Err(SumcheckError::PrismalinearRoundPolyLen { len: s.len() });
}
transcript.observe_ext(s[0]);
round_polys_eval.push(s);
let r_round = transcript.sample_ext();
debug!(%round, %r_round);
r.push(r_round);
current_evals = fold_mle_evals(current_evals, r_round);
}
if r.len() != n + 1 {
return Err(SumcheckError::PrismalinearRLen {
r_len: r.len(),
expected: n + 1,
});
}
if current_evals.values.len() != 1 {
return Err(SumcheckError::PrismalinearFinalEvalLen {
len: current_evals.values.len(),
});
}
let eval_claim = current_evals.values[0];
transcript.observe_ext(eval_claim);
Ok((
SumcheckPrismProof {
sum_claim,
s_0: s_0_ext,
round_polys_eval,
eval_claim,
},
r,
))
}