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}