use std::array::from_fn;
use itertools::Itertools;
use p3_field::Field;
use p3_matrix::dense::RowMajorMatrix;
use p3_maybe_rayon::prelude::*;
pub fn fold_mle_evals_rm<EF: Field>(mat: RowMajorMatrix<EF>, r: EF) -> RowMajorMatrix<EF> {
let width = mat.width;
let height = mat.values.len() / width;
if height <= 1 {
return mat;
}
let new_height = height / 2;
let one_minus_r = EF::ONE - r;
let values: Vec<EF> = (0..new_height)
.into_par_iter()
.flat_map(|y| {
let row0_start = (2 * y) * width;
let row1_start = (2 * y + 1) * width;
(0..width)
.map(|j| mat.values[row0_start + j] * one_minus_r + mat.values[row1_start + j] * r)
.collect::<Vec<_>>()
})
.collect();
RowMajorMatrix::new(values, width)
}
pub fn batch_fold_mle_evals_rm<EF: Field>(
mats: Vec<RowMajorMatrix<EF>>,
r: EF,
) -> Vec<RowMajorMatrix<EF>> {
mats.into_iter().map(|m| fold_mle_evals_rm(m, r)).collect()
}
pub fn sumcheck_round_poly_evals_rm<EF, FN, const WD: usize>(
n: usize,
d: usize,
parts: &[&RowMajorMatrix<EF>],
w: FN,
) -> [Vec<EF>; WD]
where
EF: Field,
FN: Fn(EF, usize, &[Vec<EF>]) -> [EF; WD] + Sync,
{
debug_assert!(parts.iter().all(|mat| {
let h = mat.values.len() / mat.width;
h == 1 << n
}));
if n == 0 {
let row_vecs: Vec<Vec<EF>> = parts.iter().map(|mat| mat.values.to_vec()).collect();
return w(EF::ONE, 0, &row_vecs).map(|x| vec![x; d]);
}
let hypercube_dim = n - 1;
let evals = (0..1usize << hypercube_dim).into_par_iter().map(|y| {
(1..=d)
.map(|x| {
let x_ef = EF::from_usize(x);
let interp_rows: Vec<Vec<EF>> = parts
.iter()
.map(|mat| {
let width = mat.width;
let r0 = (y << 1) * width;
let r1 = ((y << 1) | 1) * width;
(0..width)
.map(|j| {
let t_0 = mat.values[r0 + j];
let t_1 = mat.values[r1 + j];
t_0 + (t_1 - t_0) * x_ef
})
.collect()
})
.collect();
w(x_ef, y, &interp_rows)
})
.collect_vec()
});
let hypercube_sum = |mut acc: Vec<[EF; WD]>, x: Vec<[EF; WD]>| {
for (a, b) in acc.iter_mut().zip(x.iter()) {
for (ai, bi) in a.iter_mut().zip(b.iter()) {
*ai += *bi;
}
}
acc
};
cfg_if::cfg_if! {
if #[cfg(feature = "parallel")] {
let evals = evals.reduce(
|| vec![[EF::ZERO; WD]; d],
hypercube_sum,
);
} else {
let evals = evals.collect_vec();
let evals = evals.into_iter().fold(
vec![[EF::ZERO; WD]; d],
hypercube_sum,
);
}
}
from_fn(|i| evals.iter().map(|eval| eval[i]).collect_vec())
}
#[cfg(test)]
mod tests {
use openvm_stark_sdk::config::baby_bear_poseidon2::F;
use p3_field::PrimeCharacteristicRing;
use super::*;
#[test]
fn test_fold_mle_evals_rm_identity_at_zero() {
let mat = RowMajorMatrix::new(
vec![
F::from_u32(1),
F::from_u32(2),
F::from_u32(3),
F::from_u32(4),
F::from_u32(5),
F::from_u32(6),
F::from_u32(7),
F::from_u32(8),
],
2,
);
let folded = fold_mle_evals_rm(mat, F::ZERO);
assert_eq!(folded.width, 2);
let h = folded.values.len() / folded.width;
assert_eq!(h, 2);
assert_eq!(folded.values[0], F::from_u32(1));
assert_eq!(folded.values[1], F::from_u32(2));
assert_eq!(folded.values[2], F::from_u32(5));
assert_eq!(folded.values[3], F::from_u32(6));
}
#[test]
fn test_fold_mle_evals_rm_identity_at_one() {
let mat = RowMajorMatrix::new(
vec![
F::from_u32(1),
F::from_u32(2),
F::from_u32(3),
F::from_u32(4),
F::from_u32(5),
F::from_u32(6),
F::from_u32(7),
F::from_u32(8),
],
2,
);
let folded = fold_mle_evals_rm(mat, F::ONE);
assert_eq!(folded.values[0], F::from_u32(3));
assert_eq!(folded.values[1], F::from_u32(4));
assert_eq!(folded.values[2], F::from_u32(7));
assert_eq!(folded.values[3], F::from_u32(8));
}
#[test]
fn test_fold_mle_evals_rm_height_one() {
let mat = RowMajorMatrix::new(vec![F::from_u32(42), F::from_u32(7)], 2);
let folded = fold_mle_evals_rm(mat.clone(), F::from_u32(123));
assert_eq!(folded.values, mat.values);
}
#[test]
fn test_fold_mle_evals_rm_interpolation() {
let a = F::from_u32(10);
let b = F::from_u32(30);
let r = F::from_u32(3); let mat = RowMajorMatrix::new(vec![a, b], 1);
let folded = fold_mle_evals_rm(mat, r);
assert_eq!(folded.values.len(), 1);
let expected = a * (F::ONE - r) + b * r;
assert_eq!(folded.values[0], expected);
}
#[test]
fn test_batch_fold_mle_evals_rm() {
let mat1 = RowMajorMatrix::new(
vec![
F::from_u32(1),
F::from_u32(2),
F::from_u32(3),
F::from_u32(4),
],
2,
);
let mat2 = RowMajorMatrix::new(
vec![
F::from_u32(10),
F::from_u32(20),
F::from_u32(30),
F::from_u32(40),
],
2,
);
let r = F::ZERO;
let folded = batch_fold_mle_evals_rm(vec![mat1, mat2], r);
assert_eq!(folded.len(), 2);
assert_eq!(folded[0].values, vec![F::from_u32(1), F::from_u32(2)]);
assert_eq!(folded[1].values, vec![F::from_u32(10), F::from_u32(20)]);
}
#[test]
fn test_sumcheck_round_poly_evals_rm_identity() {
let a = F::from_u32(5);
let b = F::from_u32(11);
let mat = RowMajorMatrix::new(vec![a, b], 1);
let [result] =
sumcheck_round_poly_evals_rm::<_, _, 1>(1, 1, &[&mat], |_x, _y, parts| [parts[0][0]]);
assert_eq!(result.len(), 1);
assert_eq!(result[0], b); }
#[test]
fn test_sumcheck_round_poly_evals_rm_n2() {
let a = F::from_u32(1);
let b = F::from_u32(2);
let c = F::from_u32(3);
let d = F::from_u32(4);
let mat = RowMajorMatrix::new(vec![a, b, c, d], 1);
let [result] =
sumcheck_round_poly_evals_rm::<_, _, 1>(2, 1, &[&mat], |_x, _y, parts| [parts[0][0]]);
assert_eq!(result.len(), 1);
assert_eq!(result[0], b + d);
}
#[test]
fn test_sumcheck_round_poly_evals_rm_n0() {
let mat = RowMajorMatrix::new(vec![F::from_u32(42)], 1);
let [result] =
sumcheck_round_poly_evals_rm::<_, _, 1>(0, 1, &[&mat], |_x, _y, parts| [parts[0][0]]);
assert_eq!(result.len(), 1);
assert_eq!(result[0], F::from_u32(42));
}
#[test]
fn test_sumcheck_round_poly_evals_rm_multi_part() {
let mat1 = RowMajorMatrix::new(vec![F::from_u32(10), F::from_u32(20)], 1);
let mat2 = RowMajorMatrix::new(vec![F::from_u32(3), F::from_u32(7)], 1);
let [result] =
sumcheck_round_poly_evals_rm::<_, _, 1>(1, 1, &[&mat1, &mat2], |_x, _y, parts| {
[parts[0][0] + parts[1][0]]
});
assert_eq!(result[0], F::from_u32(27));
}
#[test]
fn test_sumcheck_round_poly_evals_rm_degree2() {
let a = F::from_u32(5);
let b = F::from_u32(11);
let mat = RowMajorMatrix::new(vec![a, b], 1);
let [result] =
sumcheck_round_poly_evals_rm::<_, _, 1>(1, 2, &[&mat], |_x, _y, parts| [parts[0][0]]);
assert_eq!(result.len(), 2);
assert_eq!(result[0], b); let expected_s2 = a * (F::ONE - F::from_u32(2)) + b * F::from_u32(2);
assert_eq!(result[1], expected_s2); }
}