binius_ip_prover/sumcheck/
bivariate_product_mle.rs1use binius_compute::Allocator;
4use binius_field::{Field, PackedField};
5use binius_math::FieldVec;
6
7use super::{
8 mle_store::MleStore,
9 quadratic_mle_evaluator::{QuadraticMleEvaluator, quadratic_mlecheck_prover},
10 round_evaluator::SharedMleCheckProver,
11};
12use crate::sumcheck::common::MleCheckProver;
13
14pub fn new<'alloc, A, F, P>(
58 alloc: &'alloc A,
59 multilinears: [FieldVec<P, A>; 2],
60 eval_point: Vec<F>,
61 eval_claim: F,
62) -> impl MleCheckProver<F> + 'alloc
63where
64 A: Allocator,
65 F: Field,
66 P: PackedField<Scalar = F>,
67{
68 quadratic_mlecheck_prover(
71 alloc,
72 multilinears,
73 |[a, b]| a * b,
74 |[a, b]| a * b,
75 eval_point,
76 eval_claim,
77 )
78}
79
80pub fn new_split_half<'alloc, A, F, P>(
94 alloc: &'alloc A,
95 buffer: FieldVec<P, A>,
96 eval_point: Vec<F>,
97 eval_claim: F,
98) -> impl MleCheckProver<F> + 'alloc
99where
100 A: Allocator,
101 F: Field,
102 P: PackedField<Scalar = F>,
103{
104 let mut store = MleStore::new(eval_point.len(), alloc);
105 let cols = store.push_split_half(buffer);
107 let evaluator =
108 QuadraticMleEvaluator::new(cols, |[a, b]: [P; 2]| a * b, |[a, b]: [P; 2]| a * b);
109 SharedMleCheckProver::new(store, [(eval_claim, evaluator)], eval_point)
110}
111
112#[cfg(test)]
113mod tests {
114 use binius_field::arch::{OptimalB128, OptimalPackedB128};
115 use binius_ip::{mlecheck, sumcheck::verify};
116 use binius_math::{
117 FieldBuffer,
118 multilinear::{eq::eq_ind, evaluate::evaluate},
119 test_utils::{random_field_buffer, random_scalars},
120 };
121 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
122
123 type StdChallenger = HasherChallenger<sha2::Sha256>;
124 use binius_compute::GlobalAllocator;
125 use itertools::{self, Itertools};
126 use rand::prelude::*;
127
128 use super::*;
129 use crate::sumcheck::{MleToSumCheckDecorator, prove::prove_single, prove_single_mlecheck};
130
131 fn test_mlecheck_prove_verify<F, P>(
132 prover: impl MleCheckProver<F>,
133 eval_claim: F,
134 eval_point: &[F],
135 multilinear_a: &FieldBuffer<P>,
136 multilinear_b: &FieldBuffer<P>,
137 ) where
138 F: Field,
139 P: PackedField<Scalar = F>,
140 {
141 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
143 let output = prove_single_mlecheck(prover, &mut prover_transcript);
144
145 prover_transcript
147 .message()
148 .write_slice(&output.multilinear_evals);
149
150 let mut verifier_transcript = prover_transcript.into_verifier();
152 let sumcheck_output = mlecheck::verify(
153 eval_point,
154 2, eval_claim,
156 &mut verifier_transcript,
157 )
158 .unwrap();
159
160 let mut reduced_eval_point = sumcheck_output.challenges.clone();
161 reduced_eval_point.reverse();
162
163 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
165
166 assert_eq!(
168 multilinear_evals[0] * multilinear_evals[1],
169 sumcheck_output.eval,
170 "Product of multilinear evaluations should equal the reduced evaluation"
171 );
172
173 let eval_a = evaluate(multilinear_a, &reduced_eval_point);
177 let eval_b = evaluate(multilinear_b, &reduced_eval_point);
178
179 assert_eq!(
180 eval_a, multilinear_evals[0],
181 "Multilinear A should evaluate to the first claimed evaluation"
182 );
183 assert_eq!(
184 eval_b, multilinear_evals[1],
185 "Multilinear B should evaluate to the second claimed evaluation"
186 );
187
188 assert_eq!(
190 output.challenges, sumcheck_output.challenges,
191 "Prover and verifier challenges should match"
192 );
193 }
194
195 fn test_wrapped_sumcheck_prove_verify<F, P>(
196 mlecheck_prover: impl MleCheckProver<F>,
197 eval_claim: F,
198 eval_point: &[F],
199 multilinear_a: &FieldBuffer<P>,
200 multilinear_b: &FieldBuffer<P>,
201 ) where
202 F: Field,
203 P: PackedField<Scalar = F>,
204 {
205 let n_vars = mlecheck_prover.n_vars();
206 let prover = MleToSumCheckDecorator::new(mlecheck_prover);
207
208 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
210 let output = prove_single(prover, &mut prover_transcript);
211
212 prover_transcript
214 .message()
215 .write_slice(&output.multilinear_evals);
216
217 let mut verifier_transcript = prover_transcript.into_verifier();
219 let sumcheck_output = verify(
220 n_vars,
221 3, eval_claim,
223 &mut verifier_transcript,
224 )
225 .unwrap();
226
227 let mut reduced_eval_point = sumcheck_output.challenges.clone();
230 reduced_eval_point.reverse();
231
232 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
234
235 let eq_ind_eval = eq_ind(eval_point, &reduced_eval_point);
237
238 assert_eq!(
240 multilinear_evals[0] * multilinear_evals[1] * eq_ind_eval,
241 sumcheck_output.eval,
242 "Product of multilinear evaluations should equal the reduced evaluation"
243 );
244
245 let eval_a = evaluate(multilinear_a, &reduced_eval_point);
248 let eval_b = evaluate(multilinear_b, &reduced_eval_point);
249
250 assert_eq!(
251 eval_a, multilinear_evals[0],
252 "Multilinear A should evaluate to the first claimed evaluation"
253 );
254 assert_eq!(
255 eval_b, multilinear_evals[1],
256 "Multilinear B should evaluate to the second claimed evaluation"
257 );
258
259 assert_eq!(
261 output.challenges, sumcheck_output.challenges,
262 "Prover and verifier challenges should match"
263 );
264 }
265
266 #[test]
267 fn test_bivariate_product_mlecheck() {
268 type F = OptimalB128;
269 type P = OptimalPackedB128;
270
271 let n_vars = 8;
272 let mut rng = StdRng::seed_from_u64(0);
273 let alloc = GlobalAllocator;
274
275 let multilinear_a = random_field_buffer::<P>(&mut rng, n_vars);
277 let multilinear_b = random_field_buffer::<P>(&mut rng, n_vars);
278
279 let product = itertools::zip_eq(multilinear_a.as_ref(), multilinear_b.as_ref())
281 .map(|(&l, &r)| l * r)
282 .collect_vec();
283 let product_buffer = FieldBuffer::new(n_vars, product);
284
285 let eval_point = random_scalars::<F>(&mut rng, n_vars);
286 let eval_claim = evaluate(&product_buffer, &eval_point);
287
288 let mlecheck_prover = new(
290 &alloc,
291 [multilinear_a.clone(), multilinear_b.clone()],
292 eval_point.clone(),
293 eval_claim,
294 );
295
296 test_mlecheck_prove_verify(
297 mlecheck_prover,
298 eval_claim,
299 &eval_point,
300 &multilinear_a,
301 &multilinear_b,
302 );
303
304 let mlecheck_prover = new(
306 &alloc,
307 [multilinear_a.clone(), multilinear_b.clone()],
308 eval_point.clone(),
309 eval_claim,
310 );
311
312 test_wrapped_sumcheck_prove_verify(
313 mlecheck_prover,
314 eval_claim,
315 &eval_point,
316 &multilinear_a,
317 &multilinear_b,
318 );
319 }
320}