binius_ip_prover/sumcheck/
bivariate_product_evaluator.rs1use binius_compute::Allocator;
4use binius_field::{Field, PackedField, WideMul};
5use binius_ip::sumcheck::RoundCoeffs;
6use binius_math::FieldVec;
7use itertools::izip;
8
9use super::{
10 mle_store::{ColId, EvaluationChunk, MleStore, RoundContext},
11 round_evals::RoundEvals,
12 round_evaluator::{SharedSumcheckProver, SumcheckRoundEvaluator},
13};
14
15pub struct BivariateProductEvaluator {
21 cols: [ColId; 2],
22}
23
24impl BivariateProductEvaluator {
25 pub const fn new(cols: [ColId; 2]) -> Self {
29 Self { cols }
30 }
31}
32
33pub fn bivariate_product_prover<'alloc, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
41 alloc: &'alloc A,
42 multilinears: [FieldVec<P, A>; 2],
43 sum: F,
44) -> SharedSumcheckProver<'alloc, A, P, BivariateProductEvaluator> {
45 assert_eq!(
46 multilinears[0].log_len(),
47 multilinears[1].log_len(),
48 "multilinears must have equal number of variables"
49 );
50
51 let mut store = MleStore::new(multilinears[0].log_len(), alloc);
52 let cols = multilinears.map(|col| store.push_owned(col));
53 SharedSumcheckProver::new(store, [(sum, BivariateProductEvaluator::new(cols))])
54}
55
56impl<F: Field, P: PackedField<Scalar = F>> SumcheckRoundEvaluator<F, P>
57 for BivariateProductEvaluator
58{
59 fn degree(&self) -> usize {
60 2
62 }
63
64 fn accumulate(&self, chunk: &EvaluationChunk<'_, P>, accum: &mut [<P as WideMul>::Output]) {
65 let a = chunk.col(self.cols[0]);
66 let b = chunk.col(self.cols[1]);
67
68 let mut y_1 = <P as WideMul>::Output::default();
74 let mut y_inf = <P as WideMul>::Output::default();
75 for (&a_0_i, &a_1_i, &b_0_i, &b_1_i) in
76 izip!(a.lo.as_ref(), a.hi.as_ref(), b.lo.as_ref(), b.hi.as_ref())
77 {
78 let a_inf_i = a_0_i + a_1_i;
80 let b_inf_i = b_0_i + b_1_i;
81
82 y_1 += P::wide_mul(a_1_i, b_1_i);
83 y_inf += P::wide_mul(a_inf_i, b_inf_i);
84 }
85
86 RoundEvals([y_1, y_inf]).add_to(accum);
87 }
88
89 fn interpolate(&self, ctx: &RoundContext<'_, P>, accum: &[P], claim: F) -> RoundCoeffs<F> {
90 let n_vars_remaining = ctx.n_vars();
92 assert!(n_vars_remaining > 0);
93
94 RoundEvals::<P, 2>::from_slots(accum)
97 .sum_scalars(n_vars_remaining)
98 .interpolate(claim)
99 }
100}
101
102#[cfg(test)]
103mod tests {
104 use binius_compute::GlobalAllocator;
105 use binius_field::arch::{OptimalB128, OptimalPackedB128};
106 use binius_ip::sumcheck::verify;
107 use binius_math::{
108 inner_product::inner_product_par, multilinear::evaluate::evaluate,
109 test_utils::random_field_buffer,
110 };
111 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
112 use rand::prelude::*;
113
114 use super::*;
115 use crate::sumcheck::prove::prove_single;
116
117 type StdChallenger = HasherChallenger<sha2::Sha256>;
118
119 #[test]
123 fn test_bivariate_product_sumcheck() {
124 type F = OptimalB128;
125 type P = OptimalPackedB128;
126
127 let n_vars = 8;
128 let mut rng = StdRng::seed_from_u64(0);
129 let alloc = GlobalAllocator;
130
131 let multilinear_a = random_field_buffer::<P>(&mut rng, n_vars);
132 let multilinear_b = random_field_buffer::<P>(&mut rng, n_vars);
133 let expected_sum = inner_product_par(&multilinear_a, &multilinear_b);
134
135 let prover = bivariate_product_prover(
136 &alloc,
137 [multilinear_a.clone(), multilinear_b.clone()],
138 expected_sum,
139 );
140
141 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
142 let output = prove_single(prover, &mut prover_transcript);
143 prover_transcript
144 .message()
145 .write_slice(&output.multilinear_evals);
146
147 let mut verifier_transcript = prover_transcript.into_verifier();
148 let sumcheck_output = verify(n_vars, 2, expected_sum, &mut verifier_transcript).unwrap();
149 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
150
151 assert_eq!(
152 multilinear_evals[0] * multilinear_evals[1],
153 sumcheck_output.eval,
154 "product of the multilinear evaluations should equal the reduced evaluation"
155 );
156
157 let mut eval_point = sumcheck_output.challenges.clone();
159 eval_point.reverse();
160 assert_eq!(evaluate(&multilinear_a, &eval_point), multilinear_evals[0]);
161 assert_eq!(evaluate(&multilinear_b, &eval_point), multilinear_evals[1]);
162 assert_eq!(output.challenges, sumcheck_output.challenges);
163 }
164}