Skip to main content

binius_ip_prover/sumcheck/
round_evaluator.rs

1// Copyright 2026 The Binius Developers
2
3//! Round evaluators over a shared [`MleStore`] and the provers that drive them.
4//!
5//! A round evaluator holds the per-round-polynomial logic for one composite claim over store
6//! columns. Evaluators hold [`ColId`](super::mle_store::ColId)s and receive column data by
7//! argument; they hold no mutable
8//! per-round state — they neither fold nor track the round claim. The driving prover owns the
9//! `RoundState` machine (the claim ↔ coeffs alternation) for each evaluator, and the store folds
10//! the columns (see the [`mle_store`](super::mle_store) module documentation).
11//!
12//! There are two evaluator traits, one per protocol:
13//! - [`SumcheckRoundEvaluator`], driven by [`SharedSumcheckProver`], emits regular sumcheck round
14//!   polynomials.
15//! - [`MleCheckRoundEvaluator`], driven by [`SharedMleCheckProver`], emits the eq-factored prime
16//!   round polynomials of the MLE-check protocol. Its claims all share one evaluation point, so the
17//!   prover — not the evaluator — owns that point's eq-indicator tracker and passes the round's eq
18//!   chunk and coordinate to the evaluator.
19//!
20//! Both provers adapt a store plus a list of evaluators — one per claim — to the [`SumcheckProver`]
21//! interface (and [`SharedMleCheckProver`] additionally to [`MleCheckProver`]), so they batch
22//! alongside standalone provers unchanged.
23
24use std::iter;
25
26use auto_impl::auto_impl;
27use binius_compute::Allocator;
28use binius_field::{Field, PackedField, WideMul};
29use binius_ip::sumcheck::RoundCoeffs;
30use binius_math::{FieldSlice, multilinear::eq::eq_ind_partial_eval};
31
32use super::{
33	MleToSumCheckEvaluator,
34	common::{MleCheckProver, SumcheckProver},
35	mle_store::{EvaluationChunk, MleStore, RoundContext},
36	round_state::RoundState,
37};
38
39/// Per-round-polynomial logic for one plain sumcheck claim over store columns.
40///
41/// The driving [`SharedSumcheckProver`] makes one parallel pass over column chunks per round; the
42/// hot loops stay monomorphized inside each evaluator and only the per-chunk [`Self::accumulate`]
43/// entry is virtual. Within a round the calls are: [`Self::accumulate`] from parallel workers into
44/// per-worker accumulator slices, then [`Self::interpolate`] once on the slot-wise summed slice.
45/// The driving prover, not the evaluator, holds the round claim and reduces the round polynomial
46/// against the verifier challenge.
47///
48/// The driving prover owns the accumulator and sizes it from the degree.
49/// - Accumulation writes wide, unreduced slots, so a per-chunk sum costs no reduction.
50/// - The prover sums the workers' slices slot-wise, then reduces once per round.
51/// - Interpolation reads the reduced slots.
52///
53/// An evaluator implements neither the allocation nor the merging.
54/// It only writes its own slots and interpolates them.
55/// The slot layout within one evaluator's run is private to it.
56///
57/// Evaluators are stateless across rounds: they hold no round claim and never fold. The prover
58/// passes the round claim into [`Self::interpolate`] and recovers it back out of the emitted round
59/// polynomial itself (as $R(0) + R(1)$), so the whole claim ↔ coeffs state machine lives in the
60/// prover's `RoundState`.
61///
62/// The `auto_impl(Box)` derive forwards the trait through `Box`, so a heterogeneous group of
63/// evaluators can drive a shared prover as `Vec<Box<dyn SumcheckRoundEvaluator<F, P>>>` while a
64/// homogeneous group avoids boxing.
65#[auto_impl(Box)]
66pub trait SumcheckRoundEvaluator<F: Field, P: PackedField<Scalar = F>>: Send + Sync {
67	/// The number of accumulator slots this evaluator's claim uses.
68	///
69	/// This is the count of sampled round-polynomial evaluations the accumulation pass collects;
70	/// the remaining evaluation is recovered from the round's sum claim in [`Self::interpolate`].
71	/// The driving prover reserves this many slots for the evaluator. It is the degree of the
72	/// accumulated (prime/composite) polynomial, which for an eq-factored MLE-check evaluator is
73	/// the prime degree, not the emitted round-polynomial degree.
74	fn degree(&self) -> usize;
75
76	/// Accumulates one chunk of the halved hypercube into `accum`.
77	///
78	/// The driving prover prepares `chunk` — the split, per-chunk column halves and eq-indicator
79	/// expansions — so the evaluator only reads its columns by [`ColId`](super::mle_store::ColId)
80	/// and eq trackers by
81	/// [`EqId`](super::mle_store::EqId). `accum` is this evaluator's run of [`Self::degree`] wide
82	/// slots, zero-initialized
83	/// on the first chunk and carried across the worker's chunks.
84	fn accumulate(&self, chunk: &EvaluationChunk<'_, P>, accum: &mut [<P as WideMul>::Output]);
85
86	/// Interpolates this round's polynomial from the accumulator and the round claim.
87	///
88	/// # Arguments
89	///
90	/// * `ctx` - This round's scalar state: the unbound-variable count, and a registered tracker's
91	///   equality coordinate and prefix.
92	/// * `accum` - This evaluator's slots, summed across every worker and reduced.
93	/// * `claim` - This evaluator's round claim.
94	fn interpolate(&self, ctx: &RoundContext<'_, P>, accum: &[P], claim: F) -> RoundCoeffs<F>;
95}
96
97/// Per-round-polynomial logic for one MLE-check claim over store columns.
98///
99/// This is the MLE-check counterpart of [`SumcheckRoundEvaluator`].
100/// It emits the prime, equality-factored round polynomials of the MLE-check protocol.
101///
102/// Every claim of one such prover shares a single evaluation point.
103/// The prover, not the evaluator, owns that point's equality indicator.
104/// Two arguments follow from that:
105/// - Accumulation receives the round's equality-indicator chunk.
106/// - Interpolation receives the round's equality coordinate.
107///
108/// So the evaluator stores no tracker identifier and no copy of the point.
109/// The accumulator contract is the one above: wide slots in, reduced slots out.
110#[auto_impl(Box)]
111pub trait MleCheckRoundEvaluator<F: Field, P: PackedField<Scalar = F>>: Send + Sync {
112	/// The number of accumulator slots this evaluator's claim uses. See
113	/// [`SumcheckRoundEvaluator::degree`].
114	fn degree(&self) -> usize;
115
116	/// Accumulates one chunk of the halved hypercube into `accum`.
117	///
118	/// `eq_ind` is the round's eq-indicator chunk (the prover looks it up once and hands it to
119	/// every evaluator); the evaluator weights its composition by it. Otherwise as
120	/// [`SumcheckRoundEvaluator::accumulate`].
121	fn accumulate(
122		&self,
123		chunk: &EvaluationChunk<'_, P>,
124		eq_ind: FieldSlice<'_, P>,
125		accum: &mut [<P as WideMul>::Output],
126	);
127
128	/// Interpolates this round's prime polynomial from the accumulator and the round claim.
129	///
130	/// As [`SumcheckRoundEvaluator::interpolate`], with one difference.
131	/// The round's equality coordinate arrives as an argument rather than from a tracker.
132	/// So `ctx` supplies only the unbound-variable count.
133	fn interpolate(
134		&self,
135		ctx: &RoundContext<'_, P>,
136		accum: &[P],
137		claim: F,
138		alpha: F,
139	) -> RoundCoeffs<F>;
140}
141
142/// Maximum log2 chunk size of the parallel round pass.
143///
144/// Chunked accumulation keeps the equality-indicator chunk resident while all evaluators read it,
145/// mirroring the chunking of the pre-store quadratic prover.
146///
147/// A chunk is 64 KiB at 128-bit scalars.
148/// A group of `N` claims reads `2N + 1` of them, so the working set is second-level, not first.
149const MAX_CHUNK_VARS: usize = 12;
150
151/// The state a group of claims shares, whichever protocol drives them.
152///
153/// One store of columns, one evaluator per claim, and each claim's position in the round cycle.
154/// Each prover below is this, plus the round pass its protocol needs.
155struct EvaluatorGroup<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>, Evaluator> {
156	store: MleStore<'a, A, P>,
157	evaluators: Vec<Evaluator>,
158	/// Each claim's position in the claim-to-coefficients cycle, parallel to the evaluators.
159	///
160	/// A claim is held directly until its round polynomial is produced.
161	/// The coefficients are held afterwards, until the fold reduces them back to a claim.
162	round_states: Vec<RoundState<RoundCoeffs<F>, F>>,
163	/// A fold challenge whose store fold waits for the next round's read pass.
164	///
165	/// Deferring it lets the columns be touched once per round instead of twice.
166	/// Only the store's fold waits here, as the round claims advance as soon as the challenge
167	/// lands.
168	buffered_challenge: Option<F>,
169}
170
171impl<'a, A, F, P, Evaluator> EvaluatorGroup<'a, A, F, P, Evaluator>
172where
173	A: Allocator,
174	F: Field,
175	P: PackedField<Scalar = F>,
176{
177	/// Creates a group from a store and the evaluators reading it, each with its initial claim.
178	fn new(
179		store: MleStore<'a, A, P>,
180		claims_with_evaluators: impl IntoIterator<Item = (F, Evaluator)>,
181	) -> Self {
182		// Unzipping lets a caller pass an array of pairs rather than two parallel vectors.
183		let (round_states, evaluators) = claims_with_evaluators
184			.into_iter()
185			.map(|(claim, evaluator)| (RoundState::Claim(claim), evaluator))
186			.unzip();
187		Self {
188			store,
189			evaluators,
190			round_states,
191			buffered_challenge: None,
192		}
193	}
194
195	/// Returns the number of variables still to bind.
196	///
197	/// A buffered challenge is a fold the store has not seen yet.
198	/// So the count sits one below the store's until the next round applies it.
199	const fn n_vars(&self) -> usize {
200		self.store.n_vars() - self.buffered_challenge.is_some() as usize
201	}
202
203	/// Adds one more claim, reading the same store, after the existing ones.
204	fn push_claim(&mut self, claim: F, evaluator: Evaluator) {
205		self.evaluators.push(evaluator);
206		self.round_states.push(RoundState::Claim(claim));
207	}
208
209	/// Interpolates every claim's round polynomial, then records it for the coming fold.
210	///
211	/// The group serves both evaluator traits, so each protocol passes its own accessors in.
212	///
213	/// # Arguments
214	///
215	/// * `accum` - Every evaluator's slots, summed across workers and reduced.
216	/// * `degree` - One evaluator's slot count.
217	/// * `interpolate` - The protocol's call into one evaluator.
218	fn interpolate_round(
219		&mut self,
220		accum: &[P],
221		degree: impl Fn(&Evaluator) -> usize,
222		interpolate: impl Fn(&Evaluator, &RoundContext<'_, P>, &[P], F) -> RoundCoeffs<F>,
223	) -> Vec<RoundCoeffs<F>> {
224		// The store has not folded this round, so its view carries this round's coordinates.
225		let ctx = self.store.round_context();
226		// Evaluators own consecutive runs of the accumulator, so one walk hands out every run.
227		let mut rest = accum;
228		let mut round_coeffs = Vec::with_capacity(self.evaluators.len());
229		for (evaluator, state) in iter::zip(&self.evaluators, &self.round_states) {
230			let (slots, tail) = rest.split_at(degree(evaluator));
231			rest = tail;
232			round_coeffs.push(interpolate(evaluator, &ctx, slots, *state.claim()));
233		}
234		debug_assert!(rest.is_empty(), "the runs must tile the accumulator exactly");
235		// The coefficients become the round state the coming fold reduces.
236		for (state, coeffs) in iter::zip(&mut self.round_states, &round_coeffs) {
237			*state = RoundState::Coeffs(coeffs.clone());
238		}
239		round_coeffs
240	}
241
242	/// Reduces every round polynomial against the challenge, forming the next round's claims.
243	///
244	/// The store's own fold is deferred to the next round's read pass.
245	fn fold(&mut self, challenge: F) {
246		for state in &mut self.round_states {
247			let claim = state.coeffs().evaluate(&challenge);
248			*state = RoundState::Claim(claim);
249		}
250		debug_assert!(
251			self.buffered_challenge.is_none(),
252			"fold called twice without an intervening execute"
253		);
254		self.buffered_challenge = Some(challenge);
255	}
256
257	/// Applies any deferred fold, then emits every column's evaluation in push order.
258	///
259	/// The store owns each column once, so each evaluation is computed once however many claims
260	/// read that column.
261	fn finish(mut self) -> Vec<F> {
262		if let Some(challenge) = self.buffered_challenge.take() {
263			self.store.fold(challenge);
264		}
265		self.store.final_evals()
266	}
267
268	/// Rebuilds the group with every evaluator replaced, carrying the store and claims across.
269	fn map_evaluators<New>(
270		self,
271		wrap: impl FnMut(Evaluator) -> New,
272	) -> EvaluatorGroup<'a, A, F, P, New> {
273		EvaluatorGroup {
274			store: self.store,
275			evaluators: self.evaluators.into_iter().map(wrap).collect(),
276			round_states: self.round_states,
277			buffered_challenge: self.buffered_challenge,
278		}
279	}
280}
281
282/// A [`SumcheckProver`] over a shared [`MleStore`] and a list of [`SumcheckRoundEvaluator`]s, one
283/// per claim.
284///
285/// Each round makes one parallel pass over the store's column chunks, feeding every evaluator,
286/// and lists the round polynomials in evaluator registration order. [`Self::fold`] folds each
287/// shared column and eq tracker once. [`Self::finish`] emits each store column's evaluation once,
288/// computed a single time by the store no matter how many claims read the column.
289pub struct SharedSumcheckProver<'a, A: Allocator, P: PackedField, Evaluator> {
290	group: EvaluatorGroup<'a, A, P::Scalar, P, Evaluator>,
291}
292
293impl<'a, A, F, P, Evaluator> SharedSumcheckProver<'a, A, P, Evaluator>
294where
295	A: Allocator,
296	F: Field,
297	P: PackedField<Scalar = F>,
298	Evaluator: SumcheckRoundEvaluator<F, P>,
299{
300	/// Creates a prover from a store and the evaluators reading its columns, each paired with its
301	/// initial claim — one `(claim, evaluator)` per claim.
302	pub fn new(
303		store: MleStore<'a, A, P>,
304		claims_with_evaluators: impl IntoIterator<Item = (F, Evaluator)>,
305	) -> Self {
306		Self {
307			group: EvaluatorGroup::new(store, claims_with_evaluators),
308		}
309	}
310
311	/// Returns a shared reference to the underlying column store.
312	pub const fn store(&self) -> &MleStore<'a, A, P> {
313		&self.group.store
314	}
315
316	/// Returns an exclusive reference to the underlying column store.
317	///
318	/// Lets a caller extend the shared store with columns that a later-added evaluator reads: the
319	/// logUp* final layer pushes the table halves onto it before adding its product evaluators.
320	pub const fn store_mut(&mut self) -> &mut MleStore<'a, A, P> {
321		&mut self.group.store
322	}
323
324	/// Adds one more evaluator — a claim reading the shared store, with its initial claim — to the
325	/// group.
326	///
327	/// Its round polynomial is appended after the existing evaluators' in [`Self::execute`].
328	pub fn add_evaluator(&mut self, claim: F, evaluator: Evaluator) {
329		self.group.push_claim(claim, evaluator);
330	}
331}
332
333impl<A, F, P, Evaluator> SumcheckProver<F> for SharedSumcheckProver<'_, A, P, Evaluator>
334where
335	A: Allocator,
336	F: Field,
337	P: PackedField<Scalar = F>,
338	Evaluator: SumcheckRoundEvaluator<F, P>,
339{
340	fn n_vars(&self) -> usize {
341		self.group.n_vars()
342	}
343
344	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
345		let n_vars_remaining = self.group.n_vars();
346		assert!(n_vars_remaining > 0);
347
348		// One parallel pass over the halved hypercube feeds every evaluator, so shared columns
349		// and eq-indicator chunks are read once per round while they are cache-resident.
350		let chunk_vars = (n_vars_remaining - 1).min(MAX_CHUNK_VARS.max(P::LOG_WIDTH));
351
352		// Each evaluator owns a contiguous run of `degree` wide slots in one flat per-worker
353		// buffer, laid out in registration order.
354		let total_slots: usize = self
355			.group
356			.evaluators
357			.iter()
358			.map(|evaluator| evaluator.degree())
359			.sum();
360
361		// The store prepares one `EvaluationChunk` per chunk of the halved hypercube — the split
362		// column halves and eq-indicator expansions each evaluator reads.
363		let group = &mut self.group;
364		let buffered_challenge = group.buffered_challenge.take();
365		let evaluators = &group.evaluators;
366		let store = &mut group.store;
367		let map = |chunk: EvaluationChunk<'_, P>| {
368			let mut accum = vec![Default::default(); total_slots];
369			// One walk hands each evaluator its own run, so no offset table is built per round.
370			let mut rest = accum.as_mut_slice();
371			for evaluator in evaluators {
372				let (slots, tail) = rest.split_at_mut(evaluator.degree());
373				rest = tail;
374				evaluator.accumulate(&chunk, slots);
375			}
376			debug_assert!(rest.is_empty(), "the runs must tile the accumulator exactly");
377			accum
378		};
379		let reduce = |mut lhs: Vec<<P as WideMul>::Output>,
380		              rhs: Vec<<P as WideMul>::Output>,
381		              _level: usize| {
382			// The only merge: sum the workers' slices slot-wise, generic over every evaluator.
383			// Plain sumcheck has no eq factor, so the reduction level is unused.
384			for (dst, src) in iter::zip(&mut lhs, rhs) {
385				*dst += src;
386			}
387			lhs
388		};
389		let accum = match buffered_challenge {
390			// The previous fold deferred its store fold; apply it and this round's read in one
391			// pass.
392			Some(challenge) => store.map_reduce_with_fold(chunk_vars, challenge, map, reduce),
393			None => store.map_reduce(chunk_vars, map, reduce),
394		};
395
396		// The workers' wide sums are complete, so every slot is reduced once for the whole round.
397		// Interpolation runs on reduced values.
398		let accum = accum.into_iter().map(P::reduce).collect::<Vec<P>>();
399
400		self.group.interpolate_round(
401			&accum,
402			|evaluator| evaluator.degree(),
403			|evaluator, ctx, slots, claim| evaluator.interpolate(ctx, slots, claim),
404		)
405	}
406
407	fn fold(&mut self, challenge: F) {
408		self.group.fold(challenge);
409	}
410
411	fn finish(self) -> Vec<F> {
412		self.group.finish()
413	}
414}
415
416/// A [`MleCheckProver`] over a shared [`MleStore`] and a group of [`MleCheckRoundEvaluator`]s.
417///
418/// This is the MLE-check counterpart of [`SharedSumcheckProver`]: every evaluator's claim shares
419/// the prover's evaluation point, and the round polynomials are the eq-factored prime polynomials
420/// of the MLE-check protocol. Because the point is shared, the prover owns its eq-indicator tracker
421/// and hands the round's eq chunk and coordinate to the evaluators, which therefore store no
422/// tracker id.
423pub struct SharedMleCheckProver<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>, Evaluator> {
424	group: EvaluatorGroup<'a, A, F, P, Evaluator>,
425	/// The evaluation point every claim shares.
426	///
427	/// No eq-indicator tracker is registered for it.
428	/// Each round expands only a chunk-wide prefix of this point and folds the higher coordinates
429	/// through the eq-weighted reduction.
430	eval_point: Vec<F>,
431}
432
433impl<'a, A, F, P, Evaluator> SharedMleCheckProver<'a, A, F, P, Evaluator>
434where
435	A: Allocator,
436	F: Field,
437	P: PackedField<Scalar = F>,
438	Evaluator: MleCheckRoundEvaluator<F, P>,
439{
440	/// Creates a prover from a store, the evaluators reading its columns each paired with its
441	/// initial claim — one `(claim, evaluator)` per claim — and the evaluation point shared by all
442	/// of the evaluators' claims.
443	pub fn new(
444		store: MleStore<'a, A, P>,
445		claims_with_evaluators: impl IntoIterator<Item = (F, Evaluator)>,
446		eval_point: Vec<F>,
447	) -> Self {
448		// precondition
449		assert_eq!(
450			eval_point.len(),
451			store.n_vars(),
452			"evaluation point length must equal the store's number of variables"
453		);
454		Self {
455			group: EvaluatorGroup::new(store, claims_with_evaluators),
456			eval_point,
457		}
458	}
459
460	/// Converts this MLE-check prover into a plain [`SharedSumcheckProver`] by folding each claim's
461	/// equality factor into its emitted round polynomials — the [Gruen24] technique of
462	/// [`MleToSumCheckEvaluator`].
463	///
464	/// This lets an eq-weighted claim batch in one evaluator group alongside plain sumcheck claims
465	/// over the same store: the logUp* pushforward reduction converts its evaluation claim on `Y`,
466	/// then adds the eq-free product evaluator to it. The store and its columns carry over
467	/// untouched; only the evaluators are wrapped, sharing this prover's eq tracker.
468	///
469	/// [Gruen24]: <https://eprint.iacr.org/2024/108>
470	pub fn into_shared_sumcheck(
471		self,
472	) -> SharedSumcheckProver<'a, A, P, Box<dyn SumcheckRoundEvaluator<F, P> + 'a>>
473	where
474		Evaluator: 'a,
475	{
476		let Self {
477			mut group,
478			eval_point,
479		} = self;
480		// The plain sumcheck prover this converts to has no eq machinery, so its wrappers read the
481		// eq indicator from a store tracker. This prover kept none, so register one here for the
482		// shared point.
483		let eq_tracker = group.store.register_eq_tracker(&eval_point);
484		// Wrap each MLE-check evaluator so it emits sumcheck round polynomials, handing it the
485		// shared eq tracker the store folds. Conversion happens before proving starts, so — since
486		// no variable has been folded — the equality prefix is one and each wrapper's sumcheck
487		// claim equals the inner MLE-check claim it carries over unchanged.
488		let group = group.map_evaluators(|evaluator| {
489			Box::new(MleToSumCheckEvaluator::new(evaluator, eq_tracker))
490				as Box<dyn SumcheckRoundEvaluator<F, P> + 'a>
491		});
492		SharedSumcheckProver { group }
493	}
494}
495
496impl<A, F, P, Evaluator> MleCheckProver<F> for SharedMleCheckProver<'_, A, F, P, Evaluator>
497where
498	A: Allocator,
499	F: Field,
500	P: PackedField<Scalar = F>,
501	Evaluator: MleCheckRoundEvaluator<F, P>,
502{
503	fn n_vars(&self) -> usize {
504		self.group.n_vars()
505	}
506
507	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
508		let n_vars_remaining = self.group.n_vars();
509		assert!(n_vars_remaining > 0);
510
511		// One parallel pass over the halved hypercube feeds every evaluator; see
512		// [`SharedSumcheckProver::execute`].
513		let chunk_vars = (n_vars_remaining - 1).min(MAX_CHUNK_VARS.max(P::LOG_WIDTH));
514		let total_slots: usize = self
515			.group
516			.evaluators
517			.iter()
518			.map(|evaluator| evaluator.degree())
519			.sum();
520
521		// The eq indicator factors into a chunk part over the low `chunk_vars` coordinates and a
522		// suffix part over the higher ones. Materialize only the chunk part, shared by every chunk
523		// and evaluator; the suffix coordinates are folded in below through `reduce`. The highest
524		// remaining coordinate is this round's `alpha`, folded into the interpolation, not the sum.
525		let alpha = self.eval_point[n_vars_remaining - 1];
526		let eq_chunk = eq_ind_partial_eval::<P>(&self.eval_point[..chunk_vars]);
527
528		let eval_point = &self.eval_point;
529		let group = &mut self.group;
530		let buffered_challenge = group.buffered_challenge.take();
531		let evaluators = &group.evaluators;
532		let store = &mut group.store;
533		let map = |chunk: EvaluationChunk<'_, P>| {
534			let mut accum = vec![Default::default(); total_slots];
535			// One walk hands each evaluator its own run, so no offset table is built per round.
536			let mut rest = accum.as_mut_slice();
537			for evaluator in evaluators {
538				let (slots, tail) = rest.split_at_mut(evaluator.degree());
539				rest = tail;
540				evaluator.accumulate(&chunk, eq_chunk.as_view(), slots);
541			}
542			debug_assert!(rest.is_empty(), "the runs must tile the accumulator exactly");
543			// Reduce the wide accumulators so the eq-weighted `reduce` can linearly extrapolate on
544			// P.
545			accum.into_iter().map(P::reduce).collect::<Vec<P>>()
546		};
547		let reduce = |mut lhs: Vec<P>, rhs: Vec<P>, level: usize| {
548			// Fold the `level`-th (suffix) coordinate: eq weights the low half by `1 - z` and the
549			// high half by `z`, i.e. the linear extrapolation `lo + z * (hi - lo)`.
550			let z = P::broadcast(eval_point[level]);
551			for (lo, hi) in iter::zip(&mut lhs, rhs) {
552				*lo += z * (hi - *lo);
553			}
554			lhs
555		};
556		let accum = match buffered_challenge {
557			Some(challenge) => store.map_reduce_with_fold(chunk_vars, challenge, map, reduce),
558			None => store.map_reduce(chunk_vars, map, reduce),
559		};
560
561		self.group.interpolate_round(
562			&accum,
563			|evaluator| evaluator.degree(),
564			|evaluator, ctx, slots, claim| evaluator.interpolate(ctx, slots, claim, alpha),
565		)
566	}
567
568	fn fold(&mut self, challenge: F) {
569		self.group.fold(challenge);
570	}
571
572	fn finish(self) -> Vec<F> {
573		self.group.finish()
574	}
575
576	fn eval_point(&self) -> &[F] {
577		&self.eval_point[..self.group.n_vars()]
578	}
579}
580
581// Prove-and-verify coverage for the shared store provers: a batched fractional-addition MLE-check,
582// and the logUp* final-layer shape (eq-weighted fractional addition batched with plain product
583// claims over shared columns).
584#[cfg(test)]
585mod tests {
586	use binius_compute::GlobalAllocator;
587	use binius_field::FieldOps;
588	use binius_ip::sumcheck::{batch_verify, batch_verify_mle};
589	use binius_math::{
590		FieldBuffer,
591		inner_product::inner_product_par,
592		multilinear::{eq::eq_ind, evaluate::evaluate},
593		test_utils::{Packed128b, random_field_buffer, random_scalars},
594		univariate::evaluate_univariate,
595	};
596	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
597	use rand::prelude::*;
598
599	use super::*;
600	use crate::sumcheck::{
601		MleToSumCheckEvaluator,
602		batch::{batch_prove, batch_prove_mle},
603		bivariate_product_evaluator::BivariateProductEvaluator,
604		frac_add_mle,
605	};
606
607	type P = Packed128b;
608	type F = <P as FieldOps>::Scalar;
609	type StdChallenger = HasherChallenger<sha2::Sha256>;
610	type CompFn = fn([P; 4]) -> P;
611
612	// The fractional-addition numerator composition, as a single-claim function.
613	fn comp_num([num_a, num_b, den_a, den_b]: [P; 4]) -> P {
614		num_a * den_b + num_b * den_a
615	}
616
617	// The fractional-addition denominator composition, as a single-claim function.
618	fn comp_den([_num_a, _num_b, den_a, den_b]: [P; 4]) -> P {
619		den_a * den_b
620	}
621
622	// Split a multilinear on its highest variable into owned low and high halves.
623	fn owned_halves(buffer: &FieldBuffer<P>) -> [FieldBuffer<P>; 2] {
624		let (lo, hi) = buffer.split_half();
625		[
626			FieldBuffer::new(lo.log_len(), lo.as_ref().into()),
627			FieldBuffer::new(hi.log_len(), hi.as_ref().into()),
628		]
629	}
630
631	// Random fractional-addition instance: four columns, evaluation point, and honest claims.
632	fn frac_instance(rng: &mut StdRng, n_vars: usize) -> ([FieldBuffer<P>; 4], Vec<F>, [F; 2]) {
633		let cols: [FieldBuffer<P>; 4] =
634			std::array::from_fn(|_| random_field_buffer::<P>(&mut *rng, n_vars));
635		let eval_point = random_scalars::<F>(&mut *rng, n_vars);
636
637		// The honest claim is each composition's MLE evaluated at the point.
638		let claims = [comp_num as CompFn, comp_den as CompFn].map(|comp| {
639			let vals = (0..1usize << n_vars)
640				.map(|i| {
641					let scalars = [
642						cols[0].get(i),
643						cols[1].get(i),
644						cols[2].get(i),
645						cols[3].get(i),
646					];
647					comp(scalars.map(P::broadcast))
648						.iter()
649						.next()
650						.expect("packed field has at least one lane")
651				})
652				.collect::<Vec<_>>();
653			evaluate(&FieldBuffer::<P>::from_values(&vals), &eval_point)
654		});
655
656		(cols, eval_point, claims)
657	}
658
659	// The store + evaluator MLE-check prover for the two fractional-addition claims, borrowing the
660	// four shared columns.
661	fn new_frac_prover<'a>(
662		alloc: &'a GlobalAllocator,
663		cols: &'a [FieldBuffer<P>; 4],
664		eval_point: &[F],
665		claims: [F; 2],
666	) -> SharedMleCheckProver<'a, GlobalAllocator, F, P, Box<dyn MleCheckRoundEvaluator<F, P>>> {
667		let mut store = MleStore::new(eval_point.len(), alloc);
668		let col_ids = cols.each_ref().map(|col| store.push(col.as_view()));
669		let (num_ev, den_ev) = frac_add_mle::evaluators(col_ids);
670		let claims_with_evaluators: [(F, Box<dyn MleCheckRoundEvaluator<F, P>>); 2] =
671			[(claims[0], Box::new(num_ev)), (claims[1], Box::new(den_ev))];
672		SharedMleCheckProver::new(store, claims_with_evaluators, eval_point.to_vec())
673	}
674
675	// Prove the two fractional-addition claims through the MLE-check batch driver, then verify.
676	// The two claims share one store, so the four columns are folded and evaluated once.
677	//
678	// The 14- and 15-variable cases exceed `MAX_CHUNK_VARS`, so the eq indicator no longer fits in
679	// one chunk prefix: they exercise `SharedMleCheckProver`'s suffix folding, where the higher eq
680	// coordinates are linearly extrapolated in `reduce` — 15 in particular hits that path through
681	// both `map_reduce` (round 0) and the fused `map_reduce_with_fold` (round 1).
682	#[test]
683	fn test_shared_frac_add_prove_verify() {
684		for n_vars in [1, 2, 3, 8, 14, 15] {
685			let mut rng = StdRng::seed_from_u64(0);
686			let (cols, eval_point, claims) = frac_instance(&mut rng, n_vars);
687
688			// Prove: one shared prover carries both claims over the four columns.
689			let alloc = GlobalAllocator;
690			let mut transcript = ProverTranscript::new(StdChallenger::default());
691			let output = batch_prove_mle(
692				vec![new_frac_prover(&alloc, &cols, &eval_point, claims)],
693				&mut transcript,
694			);
695
696			// The shared prover emits the four column evaluations once, in push order.
697			assert_eq!(output.multilinear_evals.len(), 1);
698			let evals = output.multilinear_evals[0].clone();
699			assert_eq!(evals.len(), 4);
700			transcript.message().write_scalar_slice(&evals);
701
702			// Verify: quadratic prime polynomials give degree-2 MLE-check rounds.
703			let mut verifier = transcript.into_verifier();
704			let sumcheck_output = batch_verify_mle(&eval_point, 2, &claims, &mut verifier).unwrap();
705			let verified_evals: Vec<F> = verifier.message().read_vec(4).unwrap();
706			assert_eq!(evals, verified_evals, "prover and verifier column evaluations must match");
707
708			// The prover binds variables high-to-low; `evaluate` expects them low-to-high.
709			let mut point = sumcheck_output.challenges.clone();
710			point.reverse();
711
712			// Each recovered column evaluation is the column's evaluation at the challenge point.
713			for (col, &eval) in cols.iter().zip(&verified_evals) {
714				assert_eq!(evaluate(col, &point), eval);
715			}
716
717			// The reduced evaluation is the batch combination of the two compositions at the evals.
718			let packed = std::array::from_fn(|i| P::broadcast(verified_evals[i]));
719			let composed = [comp_num, comp_den]
720				.map(|comp| comp(packed).iter().next().expect("packed field has a lane"));
721			let expected = evaluate_univariate(&composed, &sumcheck_output.batch_coeff);
722			assert_eq!(expected, sumcheck_output.eval, "reduced evaluation must match the batch");
723
724			assert_eq!(output.challenges, sumcheck_output.challenges);
725		}
726	}
727
728	// The logUp* final-layer shape: an eq-weighted fractional addition (two claims) batched with
729	// two plain product claims, all sharing the pushforward halves in one store. Prove through the
730	// sumcheck batch driver, then verify the reduced evaluation against the batched compositions.
731	#[test]
732	fn test_shared_final_layer_prove_verify() {
733		for m in [1, 2, 3, 6] {
734			let mut rng = StdRng::seed_from_u64(0);
735
736			// Three parent buffers of m variables; each splits into two m-1 variable halves.
737			let pushforward = random_field_buffer::<P>(&mut rng, m);
738			let denominator = random_field_buffer::<P>(&mut rng, m);
739			let table = random_field_buffer::<P>(&mut rng, m);
740			let [y_0, y_1] = owned_halves(&pushforward);
741			let [d_0, d_1] = owned_halves(&denominator);
742			let [t_0, t_1] = owned_halves(&table);
743
744			// Fractional claims at z; product claims are the inner products of the pushforward and
745			// table halves.
746			let z = random_scalars::<F>(&mut rng, m - 1);
747			let frac_claims: [F; 2] = [comp_num as CompFn, comp_den as CompFn].map(|comp| {
748				let vals = (0..1usize << (m - 1))
749					.map(|i| {
750						let scalars = [y_0.get(i), y_1.get(i), d_0.get(i), d_1.get(i)];
751						comp(scalars.map(P::broadcast))
752							.iter()
753							.next()
754							.expect("packed field has at least one lane")
755					})
756					.collect::<Vec<_>>();
757				evaluate(&FieldBuffer::<P>::from_values(&vals), &z)
758			});
759			let e_0 = inner_product_par(&y_0, &t_0);
760			let e_1 = inner_product_par(&y_1, &t_1);
761
762			// One store with six borrowed columns, in push order [Y_0, Y_1, D_0, D_1, T_0, T_1].
763			let alloc = GlobalAllocator;
764			let mut store = MleStore::new(m - 1, &alloc);
765			let [y_0_col, y_1_col, d_0_col, d_1_col, t_0_col, t_1_col] =
766				[&y_0, &y_1, &d_0, &d_1, &t_0, &t_1].map(|col| store.push(col.as_view()));
767
768			// The eq-weighted fractional evaluators, wrapped so they emit sumcheck round
769			// polynomials. The wrappers, driven by a plain sumcheck prover, hold the shared eq
770			// tracker themselves.
771			let (num_evaluator, den_evaluator) =
772				frac_add_mle::evaluators([y_0_col, y_1_col, d_0_col, d_1_col]);
773			let eq_tracker = store.register_eq_tracker(&z);
774			let num_evaluator = MleToSumCheckEvaluator::new(num_evaluator, eq_tracker);
775			let den_evaluator = MleToSumCheckEvaluator::new(den_evaluator, eq_tracker);
776
777			// The two plain product claims over the pushforward and table halves.
778			let product_0 = BivariateProductEvaluator::new([y_0_col, t_0_col]);
779			let product_1 = BivariateProductEvaluator::new([y_1_col, t_1_col]);
780
781			// Claims paired with evaluators in order: the two fractional claims, then the two
782			// product sums.
783			let claims_with_evaluators: [(F, Box<dyn SumcheckRoundEvaluator<F, P>>); 4] = [
784				(frac_claims[0], Box::new(num_evaluator)),
785				(frac_claims[1], Box::new(den_evaluator)),
786				(e_0, Box::new(product_0)),
787				(e_1, Box::new(product_1)),
788			];
789			let shared = SharedSumcheckProver::new(store, claims_with_evaluators);
790
791			// Prove and record the four claim sums in evaluator order.
792			let mut transcript = ProverTranscript::new(StdChallenger::default());
793			let output = batch_prove(vec![shared], &mut transcript);
794
795			// The shared prover emits each store column's evaluation once, in push order.
796			assert_eq!(output.multilinear_evals.len(), 1);
797			let evals = output.multilinear_evals[0].clone();
798			assert_eq!(evals.len(), 6);
799			transcript.message().write_scalar_slice(&evals);
800
801			// Verify: the eq-wrapped fractional rounds have degree 3, the product rounds degree 2,
802			// so the batched round polynomial has degree 3.
803			let claims = [frac_claims[0], frac_claims[1], e_0, e_1];
804			let mut verifier = transcript.into_verifier();
805			let sumcheck_output = batch_verify(m - 1, 3, &claims, &mut verifier).unwrap();
806			let verified_evals: Vec<F> = verifier.message().read_vec(6).unwrap();
807			assert_eq!(evals, verified_evals, "prover and verifier column evaluations must match");
808
809			// The prover binds variables high-to-low; `evaluate` expects them low-to-high.
810			let mut point = sumcheck_output.challenges.clone();
811			point.reverse();
812
813			// Each recovered column evaluation is the column's evaluation at the challenge point.
814			for (col, &eval) in [&y_0, &y_1, &d_0, &d_1, &t_0, &t_1]
815				.iter()
816				.zip(&verified_evals)
817			{
818				assert_eq!(evaluate(col, &point), eval);
819			}
820
821			// The reduced evaluation batches the four claims' reduced compositions in evaluator
822			// order. The fractional compositions carry the equality factor at z; the products do
823			// not.
824			let [y0, y1, d0, d1, t0, t1] =
825				<[F; 6]>::try_from(verified_evals).expect("six column evaluations");
826			let eq = eq_ind(&z, &point);
827			let reduced = [(y0 * d1 + y1 * d0) * eq, (d0 * d1) * eq, y0 * t0, y1 * t1];
828			let expected = evaluate_univariate(&reduced, &sumcheck_output.batch_coeff);
829			assert_eq!(expected, sumcheck_output.eval, "reduced evaluation must match the batch");
830
831			assert_eq!(output.challenges, sumcheck_output.challenges);
832		}
833	}
834}