Skip to main content

binius_ip_prover/sumcheck/
quadratic_mle_evaluator.rs

1// Copyright 2026 The Binius Developers
2
3use binius_compute::Allocator;
4use binius_field::{Field, PackedField, WideMul};
5use binius_ip::sumcheck::RoundCoeffs;
6use binius_math::{FieldSlice, FieldVec};
7
8use super::{
9	mle_store::{ColId, ColumnChunk, EvaluationChunk, MleStore, RoundContext},
10	round_evals::RoundEvals,
11	round_evaluator::{MleCheckRoundEvaluator, SharedMleCheckProver},
12};
13
14/// MLE-check round evaluator for one quadratic composition over N store columns.
15///
16/// This is the store-backed successor of the quadratic MLE-check prover: it evaluates the
17/// composition in one pass per round, using the degree-2 interpolation of [Gruen24] section 3.2.
18/// Batch
19/// several quadratic MLE checks by registering one evaluator per claim on a shared store; they read
20/// the shared columns from the same round pass.
21///
22/// [Gruen24]: <https://eprint.iacr.org/2024/108>
23///
24/// The evaluator emits the prime (eq-factored) round polynomial of the MLE-check protocol. Wrap it
25/// in [`MleToSumCheckEvaluator`](super::MleToSumCheckEvaluator) to emit a regular sumcheck round
26/// polynomial.
27pub struct QuadraticMleEvaluator<Composition, InfinityComposition, const N: usize> {
28	// Store columns holding the packed evaluations of the input multilinears.
29	cols: [ColId; N],
30	// Full quadratic composition evaluated on the "x = 1" branch for each multilinear.
31	composition: Composition,
32	// Composition restricted to highest-degree terms for the "x = ∞" evaluation (Karatsuba).
33	infinity_composition: InfinityComposition,
34}
35
36impl<Composition, InfinityComposition, const N: usize>
37	QuadraticMleEvaluator<Composition, InfinityComposition, N>
38{
39	/// Creates an evaluator over the store columns `cols`.
40	///
41	/// The evaluator holds no eq-indicator tracker and no copy of the evaluation point: the driving
42	/// [`SharedMleCheckProver`] owns the shared point's tracker and passes the round's eq chunk and
43	/// coordinate in. The claimed evaluation is likewise held by the prover, not the evaluator.
44	///
45	/// # Arguments
46	///
47	/// * `cols` - The N store columns the composition reads.
48	/// * `composition` - Evaluates the quadratic composition of the N column values.
49	/// * `infinity_composition` - The composition restricted to its highest-degree terms, for the
50	///   Karatsuba evaluation at infinity.
51	pub fn new(
52		cols: [ColId; N],
53		composition: Composition,
54		infinity_composition: InfinityComposition,
55	) -> Self {
56		// precondition
57		assert!(N > 0);
58
59		Self {
60			cols,
61			composition,
62			infinity_composition,
63		}
64	}
65}
66
67/// Builds an MLE-check prover for one quadratic composition of N owned multilinears.
68///
69/// The reduced claim is
70///
71/// ```text
72///     sum_{v in B} C(M_0(v), ..., M_{N-1}(v)) * eq(v, eval_point) = eval_claim
73/// ```
74///
75/// for composition `C` over the boolean hypercube `B`.
76/// It reduces to one evaluation claim per input multilinear at the challenge point.
77///
78/// This is the single-claim path.
79/// Several quadratic claims that share columns are instead proved together by registering one
80/// evaluator per claim on a shared store, which folds each shared column only once.
81///
82/// # Arguments
83///
84/// * `multilinears` - The N input multilinears, each over `eval_point.len()` variables.
85/// * `composition` - Evaluates the quadratic composition, e.g. `|[a, b, c]| a * b - c`.
86/// * `infinity_composition` - The composition's highest-degree terms, for the Karatsuba evaluation
87///   at infinity, e.g. `|[a, b, _c]| a * b`.
88/// * `eval_point` - The point at which the composite MLE is claimed.
89/// * `eval_claim` - The claimed evaluation of the composite MLE at `eval_point`.
90///
91/// # Returns
92///
93/// A prover whose reduction emits the N column evaluations in the order given.
94///
95/// # Panics
96///
97/// Panics if any multilinear's variable count differs from `eval_point.len()`, or if `N == 0`.
98pub fn quadratic_mlecheck_prover<
99	'alloc,
100	A,
101	F,
102	P,
103	Composition,
104	InfinityComposition,
105	const N: usize,
106>(
107	alloc: &'alloc A,
108	multilinears: [FieldVec<P, A>; N],
109	composition: Composition,
110	infinity_composition: InfinityComposition,
111	eval_point: Vec<F>,
112	eval_claim: F,
113) -> SharedMleCheckProver<'alloc, A, F, P, QuadraticMleEvaluator<Composition, InfinityComposition, N>>
114where
115	A: Allocator,
116	F: Field,
117	P: PackedField<Scalar = F>,
118	Composition: Fn([P; N]) -> P + Send + Sync,
119	InfinityComposition: Fn([P; N]) -> P + Send + Sync,
120{
121	let mut store = MleStore::new(eval_point.len(), alloc);
122	// Hand each column to the store, which checks its variable count against the point length.
123	let cols = multilinears.map(|col| store.push_owned(col));
124	let evaluator = QuadraticMleEvaluator::new(cols, composition, infinity_composition);
125	SharedMleCheckProver::new(store, [(eval_claim, evaluator)], eval_point)
126}
127
128impl<F, P, Composition, InfinityComposition, const N: usize> MleCheckRoundEvaluator<F, P>
129	for QuadraticMleEvaluator<Composition, InfinityComposition, N>
130where
131	F: Field,
132	P: PackedField<Scalar = F>,
133	Composition: Fn([P; N]) -> P + Send + Sync,
134	InfinityComposition: Fn([P; N]) -> P + Send + Sync,
135{
136	fn degree(&self) -> usize {
137		// Quadratic composition: two sampled evaluations, `y_1` and `y_inf`.
138		2
139	}
140
141	fn accumulate(
142		&self,
143		chunk: &EvaluationChunk<'_, P>,
144		eq_ind: FieldSlice<'_, P>,
145		accum: &mut [<P as WideMul>::Output],
146	) {
147		// Each column arrives split into low/high halves for the top variable: the low half
148		// corresponds to x=0, the high half to x=1.
149		let cols: [&ColumnChunk<'_, P>; N] = self.cols.map(|id| chunk.col(id));
150
151		// Bind each half to a slice ahead of the element loop.
152		// A half is a base pointer and a length, so deriving it per element makes it a memory read.
153		//
154		//     per element, 4 columns:  derived in-loop   9 vector + 29 scalar loads
155		//                              bound here        9 vector loads
156		let los: [&[P]; N] = cols.map(|col| col.lo.as_ref());
157		let his: [&[P]; N] = cols.map(|col| col.hi.as_ref());
158
159		let mut y_1 = <P as WideMul>::Output::default();
160		let mut y_inf = <P as WideMul>::Output::default();
161		for (idx, &eq_i) in eq_ind.as_ref().iter().enumerate() {
162			// Gather the idx-th evaluations of every multilinear at both halves.
163			let mut evals_1 = [P::default(); N];
164			let mut evals_inf = [P::default(); N];
165
166			for i in 0..N {
167				let lo_i = los[i][idx];
168				let hi_i = his[i][idx];
169
170				// Compose once with the high half and once with the lo+hi combination.
171				// The lo+hi branch corresponds to evaluation at infinity for multilinears.
172				evals_1[i] = hi_i;
173				evals_inf[i] = lo_i + hi_i;
174			}
175
176			// Weight the composition by the eq indicator to keep the sumcheck claim aligned to
177			// eval_point. Only this final multiply is widened; the composition products are already
178			// reduced.
179			y_1 += P::wide_mul((self.composition)(evals_1), eq_i);
180			y_inf += P::wide_mul((self.infinity_composition)(evals_inf), eq_i);
181		}
182
183		RoundEvals([y_1, y_inf]).add_to(accum);
184	}
185
186	fn interpolate(
187		&self,
188		ctx: &RoundContext<'_, P>,
189		accum: &[P],
190		claim: F,
191		alpha: F,
192	) -> RoundCoeffs<F> {
193		// The store has not yet folded this round, so its remaining-variable count is this round's.
194		let n_vars_remaining = ctx.n_vars();
195		assert!(n_vars_remaining > 0);
196
197		// `accum` is already reduced (the prover's `map` pass reduced the wide accumulators). Sum
198		// the packed lanes into scalars, then interpolate. `claim` is this round's prime eval;
199		// `alpha`, this round's eq coordinate, ties it to the point.
200		RoundEvals::<P, 2>::from_slots(accum)
201			.sum_scalars(n_vars_remaining)
202			.interpolate_eq(claim, alpha)
203	}
204}
205
206#[cfg(test)]
207mod tests {
208	use std::{array, iter};
209
210	use binius_compute::GlobalAllocator;
211	use binius_field::{arch::OptimalPackedB128, field::FieldOps};
212	use binius_ip::mlecheck;
213	use binius_math::{
214		FieldBuffer,
215		multilinear::evaluate::evaluate,
216		test_utils::{random_field_buffer, random_scalars},
217	};
218	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
219	use itertools::Itertools;
220	use rand::prelude::*;
221
222	use super::*;
223	use crate::sumcheck::prove_single_mlecheck;
224
225	type StdChallenger = HasherChallenger<sha2::Sha256>;
226
227	// Prove one quadratic MLE-check via the shared store, then verify it through the verifier.
228	// The verifier's reduced evaluation must equal the composition of the recovered column evals.
229	// Each column must evaluate to its claimed value at the challenge point.
230	// Prover and verifier challenges must agree.
231	fn prove_verify<F, P, const N: usize>(
232		composition: impl Fn([P; N]) -> P + Clone + Send + Sync,
233		infinity_composition: impl Fn([P; N]) -> P + Send + Sync,
234	) where
235		F: Field,
236		P: PackedField<Scalar = F>,
237	{
238		let n_vars = 8;
239		let mut rng = StdRng::seed_from_u64(0);
240		let alloc = GlobalAllocator;
241
242		let multilinears: [_; N] = array::from_fn(|_| random_field_buffer::<P>(&mut rng, n_vars));
243
244		// The honest claim is the composite MLE evaluated at the point.
245		let composite_vals = (0..1 << n_vars.saturating_sub(P::LOG_WIDTH))
246			.map(|i| composition(array::from_fn(|j| multilinears[j].as_ref()[i])))
247			.collect_vec();
248		let composite_vals = FieldBuffer::new(n_vars, composite_vals);
249		let eval_point = random_scalars::<F>(&mut rng, n_vars);
250		let eval_claim = evaluate(&composite_vals, &eval_point);
251
252		let prover = quadratic_mlecheck_prover(
253			&alloc,
254			multilinears.clone(),
255			composition.clone(),
256			infinity_composition,
257			eval_point.clone(),
258			eval_claim,
259		);
260
261		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
262		let output = prove_single_mlecheck(prover, &mut prover_transcript);
263		prover_transcript
264			.message()
265			.write_slice(&output.multilinear_evals);
266
267		let mut verifier_transcript = prover_transcript.into_verifier();
268		let sumcheck_output = mlecheck::verify(
269			&eval_point,
270			2, // quadratic compositions have degree-2 round polynomials
271			eval_claim,
272			&mut verifier_transcript,
273		)
274		.unwrap();
275
276		// The prover binds variables high-to-low.
277		// `evaluate` expects them low-to-high, so reverse the challenges.
278		let mut reduced_eval_point = sumcheck_output.challenges.clone();
279		reduced_eval_point.reverse();
280
281		let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(N).unwrap();
282
283		// The reduced evaluation is the composition of the column evaluations.
284		let evals_packed: [P; N] = array::from_fn(|i| P::broadcast(multilinear_evals[i]));
285		assert_eq!(
286			composition(evals_packed).iter().next().unwrap(),
287			sumcheck_output.eval,
288			"composition of the column evaluations should equal the reduced evaluation"
289		);
290
291		// Each column evaluates to its claimed value at the challenge point.
292		for (multilinear, claimed_eval) in iter::zip(&multilinears, multilinear_evals) {
293			assert_eq!(evaluate(multilinear, &reduced_eval_point), claimed_eval);
294		}
295
296		assert_eq!(
297			output.challenges, sumcheck_output.challenges,
298			"prover and verifier challenges should match"
299		);
300	}
301
302	#[test]
303	fn test_identity_mlecheck() {
304		// One column returned unchanged: the narrowest gather a round pass performs.
305		prove_verify::<_, OptimalPackedB128, 1>(|[a]| a, |[_a]| OptimalPackedB128::zero());
306	}
307
308	#[test]
309	fn test_linear_mlecheck() {
310		prove_verify::<_, OptimalPackedB128, 2>(
311			|[a, b]| a + b,
312			|[_a, _b]| OptimalPackedB128::zero(),
313		);
314	}
315
316	#[test]
317	fn test_bivariate_product_mlecheck() {
318		prove_verify::<_, OptimalPackedB128, 2>(|[a, b]| a * b, |[a, b]| a * b);
319	}
320
321	#[test]
322	fn test_mul_gate_mlecheck() {
323		prove_verify::<_, OptimalPackedB128, 3>(|[a, b, c]| a * b - c, |[a, b, _c]| a * b);
324	}
325
326	#[test]
327	fn test_4_variate_composition_mlecheck() {
328		prove_verify::<_, OptimalPackedB128, 4>(
329			|[a, b, c, d]| (a + b) * (c + d),
330			|[a, b, c, d]| (a + b) * (c + d),
331		);
332	}
333}