Skip to main content

binius_ip_prover/sumcheck/
bivariate_product_mle.rs

1// Copyright 2023-2025 Irreducible Inc.
2
3use 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
14/// Creates an [`MleCheckProver`] that reduces an evaluation claim on a multilinear extension
15/// of the product of two multilinears to evaluation claims on said multilinears.
16///
17/// ## Mathematical Definition
18/// * $n \in N$ - number of variables in multilinear polynomials
19/// * $A, B \in F\[x\], x = \(x_1, \ldots, x_n\)$ - multilinears being multiplied
20/// * $(\widetilde{AB})\[x\] = y$ - evaluation claim on the product MLE
21///
22/// The claim is equivalent to $P(x) = \sum_{v \in B} \widetilde{eq}(v, x) A(v) B(v) = y$, and the
23/// reduction can be achieved by sumchecking the latter degree-3 composition. The paper [Gruen24],
24/// however, describes a way to partition the $\widetilde{eq}(v, x)$ into three parts in round $j
25/// \in 1, \ldots, n$ during specialization of variable $v_{n-j+1}$, with $j-1$ challenges
26/// $\alpha_i$ already sampled:
27///
28/// $$ \widetilde{eq}(x_{n-j+2}, \ldots, x_n; \alpha_{j-1}, \ldots, \alpha_{1}) \tag{1} $$
29/// $$ \widetilde{eq}(x_{n-j+1}; v_{n-j+1}) \tag{2} $$
30/// $$ \widetilde{eq}(x_1, \ldots, x_{n-j}; v_1, \ldots, v_{n-j}) \tag{3} $$
31///
32/// The following holds:
33/// * (1) is a constant that can be incrementally updated in O(1) time,
34/// * (2) is a linear polynomial that is easy to compute in monomial form specialized to either
35///   variable
36/// * (3) is a an equality indicator over the claim point suffix
37///
38/// These observations allow us to instead sumcheck:
39/// $$
40/// P'(x) = \sum_{v \in B} \widetilde{eq}(x_1, \ldots, x_{n-j}; v_1, \ldots, v_{n-j}) A(v) B(v)
41/// $$
42///
43/// Which is simpler because:
44/// * $P'(x)$ is degree-2 in $j$-th variable, requiring one less evaluation point
45/// * Equality indicator expansion does not depend on $j$-th variable and thus doesn't need to be
46///   interpolated
47///
48/// After computing the round polynomial for $P'(x)$ in monomial form, one can simply multiply by
49/// (2) and (1) in polynomial form. For more details, see the
50/// [equality trackers](crate::sumcheck::eq_tracker) and [Gruen24] Section 3.2.
51///
52/// Note 1: as evident from the definition, this prover binds variables in high-to-low index order.
53///
54/// Note 2: evaluation points are 0 (implicit), 1 and Karatsuba infinity.
55///
56/// [Gruen24]: <https://eprint.iacr.org/2024/108>
57pub 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	// The product is symmetric, so the infinity composition (highest-degree terms) equals the full
69	// composition.
70	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
80/// Reduces the product of the two halves of a single buffer, sharing one allocation.
81///
82/// `buffer` has one more variable than `eval_point`:
83/// - its low half fixes the highest variable to 0,
84/// - its high half fixes it to 1.
85///
86/// The reduction proves the product of those two halves.
87/// Both halves live inside the one buffer, so separating them costs no copy.
88/// This is the zero-copy path the product-check layer reduction uses on its large witness layers.
89///
90/// # Returns
91///
92/// A prover whose reduction emits the low half's evaluation, then the high half's.
93pub 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	// The store checks that the buffer has exactly one more variable than itself.
106	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		// Run the proving protocol
142		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
143		let output = prove_single_mlecheck(prover, &mut prover_transcript);
144
145		// Write the multilinear evaluations to the transcript
146		prover_transcript
147			.message()
148			.write_slice(&output.multilinear_evals);
149
150		// Convert to verifier transcript and run verification
151		let mut verifier_transcript = prover_transcript.into_verifier();
152		let sumcheck_output = mlecheck::verify(
153			eval_point,
154			2, // degree 2 for bivariate product
155			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		// Read the multilinear evaluations from the transcript
164		let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
165
166		// Check that the product of the evaluations equals the reduced evaluation
167		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		// Check that the original multilinears evaluate to the claimed values at the challenge
174		// point The prover binds variables from high to low, but evaluate expects them from low
175		// to high
176		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		// Also verify the challenges match what the prover saw
189		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		// Run the proving protocol
209		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
210		let output = prove_single(prover, &mut prover_transcript);
211
212		// Write the multilinear evaluations to the transcript
213		prover_transcript
214			.message()
215			.write_slice(&output.multilinear_evals);
216
217		// Convert to verifier transcript and run verification
218		let mut verifier_transcript = prover_transcript.into_verifier();
219		let sumcheck_output = verify(
220			n_vars,
221			3, // degree 3 for trivariate product (bivariate by equality indicator)
222			eval_claim,
223			&mut verifier_transcript,
224		)
225		.unwrap();
226
227		// The prover binds variables from high to low, but evaluate expects them from low
228		// to high
229		let mut reduced_eval_point = sumcheck_output.challenges.clone();
230		reduced_eval_point.reverse();
231
232		// Read the multilinear evaluations from the transcript
233		let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
234
235		// Evaluate the equality indicator
236		let eq_ind_eval = eq_ind(eval_point, &reduced_eval_point);
237
238		// Check that the product of the evaluations equals the reduced evaluation
239		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		// Check that the original multilinears evaluate to the claimed values at the challenge
246		// point
247		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		// Also verify the challenges match what the prover saw
260		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		// Generate two random multilinear polynomials
276		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		// Compute product multilinear
280		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		// Create the prover
289		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		// Create another prover for the wrapped test
305		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}