Skip to main content

binius_ip_prover/sumcheck/
selector_mle.rs

1// Copyright 2023-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use binius_compute::BufferData;
5use binius_field::{Field, PackedField, WideMul};
6use binius_ip::sumcheck::RoundCoeffs;
7use binius_math::{FieldBuffer, multilinear::fold::fold_highest_var_inplace};
8use binius_utils::{bitwise::Bitwise, rayon::prelude::*};
9use itertools::izip;
10
11use super::{
12	common::SumcheckProver, eq_tracker::ChunkedEqTracker, round_evals::RoundEvals,
13	round_state::RoundState, switchover::BinarySwitchover,
14};
15
16pub struct Claim<F: Field> {
17	pub point: Vec<F>,
18	pub value: F,
19}
20
21/// A [`SumcheckProver`] implementation that proves an mlecheck over many compositions of the
22/// form `selected * selector + (1 - selector)`, where `selected` is the shared large field
23/// multilinear and `selector` comes from the set of 1-bit multilinears. Unlike other multi mlecheck
24/// provers however the evaluation point is _not_ shared but is specified per selector.
25///
26/// The set of 1-bit multilinears is represented by a power-of-two long slice of bitmasks, and the
27/// multilinear set is constructed by arranging the bitmasks as a 2D matrix in row-major order and
28/// taking vertical slices. This representation is very compact and has no embedding overhead.
29///
30/// To combat memory blowup issues arising from folding 1-bit multilinears, this prover introduces
31/// switchover. See `BinarySwitchover` for more in-depth explanation of the mechanism. Also note
32/// that the need to expand the equality indicator for each multilinear still results in some
33/// blowup.
34pub struct SelectorMlecheckProver<'b, P: PackedField, B: Bitwise, Data: BufferData<P> = Vec<P>> {
35	last_coeffs_or_sums: RoundState<Vec<RoundCoeffs<P::Scalar>>, Vec<P::Scalar>>,
36	selected: FieldBuffer<P, Data>,
37	eq_trackers: Vec<ChunkedEqTracker<P>>,
38	weights: Vec<P::Scalar>,
39	switchover: BinarySwitchover<'b, P, B>,
40}
41
42impl<'b, F: Field, P: PackedField<Scalar = F>, B: Bitwise, Data: BufferData<P>>
43	SelectorMlecheckProver<'b, P, B, Data>
44{
45	/// Constructs a prover, given `bitmasks` as representation of 1-bit columns, `selected` being
46	/// the shared large field multilinear, individual `claims` per selector, `weights` to combine
47	/// the per-selector round polynomials into one (one weight per claim), and `switchover` as the
48	/// round at which 1-bit columns should be folded.
49	///
50	/// The prover exposes a single claim — the `weights`-combination `Σ_i weights[i] · C_i` of the
51	/// per-selector claims. Supplying the equality-indicator tensor `eq_k(γ, ·)` as the weights
52	/// batches the claims with `eq_k(γ, i)`.
53	pub fn new(
54		selected: FieldBuffer<P, Data>,
55		claims: Vec<Claim<F>>,
56		bitmasks: &'b [B],
57		weights: Vec<F>,
58		switchover: usize,
59	) -> Self {
60		let n_vars = selected.log_len();
61
62		assert!(
63			claims.iter().all(|claim| claim.point.len() == n_vars),
64			"multilinears must have equal number of variables"
65		);
66
67		assert_eq!(
68			weights.len(),
69			claims.len(),
70			"number of weights must match the number of claims"
71		);
72
73		assert_eq!(
74			bitmasks.len(),
75			selected.len(),
76			"bitmasks slice length must match the selected multilinear length"
77		);
78
79		const MAX_CHUNK_VARS: usize = 8;
80		let (eq_trackers, sums) = claims
81			.into_par_iter()
82			.map(|Claim { point, value }| (ChunkedEqTracker::new(MAX_CHUNK_VARS, &point), value))
83			.collect::<(Vec<_>, Vec<_>)>();
84
85		let switchover = BinarySwitchover::new(sums.len(), switchover.min(n_vars), bitmasks);
86		let last_coeffs_or_sums = RoundState::Claim(sums);
87
88		Self {
89			last_coeffs_or_sums,
90			selected,
91			eq_trackers,
92			weights,
93			switchover,
94		}
95	}
96}
97
98impl<'b, F, P, B, Data> SumcheckProver<F> for SelectorMlecheckProver<'b, P, B, Data>
99where
100	F: Field,
101	P: PackedField<Scalar = F>,
102	B: Bitwise,
103	Data: BufferData<P>,
104{
105	fn n_vars(&self) -> usize {
106		self.selected.log_len()
107	}
108
109	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
110		let sums = self.last_coeffs_or_sums.claim();
111
112		assert!(self.n_vars() > 0);
113
114		// Perform chunked summation: for every row, evaluate all compositions and add up
115		// results to an array of round evals accumulators. Alternative would be to sum each
116		// composition on its own pass, but that would require reading the entirety of eq field
117		// buffer on each pass, which will evict the latter from the cache. By doing chunked
118		// compute, we reasonably hope that eq chunk always stays in L1 cache. We can also
119		// leverage the outer product representation of the eq indicator.
120		//
121		// We also do switchover there, which by definition requires small scratchpads to hold
122		// large field partial evaluations of the transparent multilinears.
123		let chunk_vars = self
124			.eq_trackers
125			.first()
126			.map(|eq_tracker| eq_tracker.chunk().log_len())
127			.unwrap_or_default();
128		let chunk_count = 1 << (self.n_vars() - 1 - chunk_vars);
129
130		// The fold below reads both halves concurrently from many rayon tasks.
131		// Borrowed halves cross that boundary for any backing store the buffer is built on.
132		let (selected_0, selected_1) = self.selected.split_half();
133
134		let packed_prime_evals = (0..chunk_count)
135			.into_par_iter()
136			.fold(
137				|| {
138					(
139						vec![RoundEvals::<P, 2>::default(); sums.len()],
140						FieldBuffer::<P>::zeros(chunk_vars),
141						FieldBuffer::<P>::zeros(chunk_vars),
142					)
143				},
144				|(mut packed_prime_evals, mut binary_chunk_0, mut binary_chunk_1), chunk_index| {
145					let selected_0_chunk = selected_0.chunk(chunk_vars, chunk_index);
146					let selected_1_chunk = selected_1.chunk(chunk_vars, chunk_index);
147
148					for (bit_offset, (round_evals, eq_tracker)) in
149						izip!(&mut packed_prime_evals, &self.eq_trackers).enumerate()
150					{
151						let eq_chunk = eq_tracker.chunk();
152						let eq_suffix_eval = eq_tracker.suffix().get(chunk_index);
153
154						let selector_0_chunk = self.switchover.get_chunk(
155							&mut binary_chunk_0,
156							bit_offset,
157							chunk_vars,
158							chunk_index,
159						);
160
161						let selector_1_chunk = self.switchover.get_chunk(
162							&mut binary_chunk_1,
163							bit_offset,
164							chunk_vars,
165							chunk_index | chunk_count,
166						);
167
168						// Accumulate `eq_i * composition` in unreduced (wide) form and reduce once
169						// at the end of the chunk. Only the final multiply by `eq_i` is widened;
170						// the `composition` product is reduced as usual because it feeds into that
171						// widening multiply.
172						let mut wide_y_1 = <P as WideMul>::Output::default();
173						let mut wide_y_inf = <P as WideMul>::Output::default();
174						for (&eq_i, &selected_0_i, &selected_1_i, &selector_0_i, &selector_1_i) in izip!(
175							eq_chunk.as_ref(),
176							selected_0_chunk.as_ref(),
177							selected_1_chunk.as_ref(),
178							selector_0_chunk.as_ref(),
179							selector_1_chunk.as_ref(),
180						) {
181							let selected_inf_i = selected_0_i + selected_1_i;
182							let selector_inf_i = selector_0_i + selector_1_i;
183
184							// selected * selector + (1 - selector)
185							// @one: selector * (selected - 1) + 1
186							// @inf: selector * selected (note that lower degree terms are dropped)
187							let y_1_prod = selector_1_i * (selected_1_i - P::one()) + P::one();
188							let y_inf_prod = selector_inf_i * selected_inf_i;
189							wide_y_1 += P::wide_mul(eq_i, y_1_prod);
190							wide_y_inf += P::wide_mul(eq_i, y_inf_prod);
191						}
192						let chunk_round_evals = RoundEvals([wide_y_1, wide_y_inf]).reduce::<P>();
193
194						// Apply the common factor from the outer product representation of the eq
195						// ind
196						*round_evals += &(chunk_round_evals * eq_suffix_eval);
197					}
198
199					(packed_prime_evals, binary_chunk_0, binary_chunk_1)
200				},
201			)
202			.map(|(evals, _, _)| evals)
203			// A merge seeded with a partial that already exists never touches a buffer of zeros.
204			// An identity would allocate and zero one accumulator per merge, then add all of it.
205			.reduce_with(|lhs, rhs| izip!(lhs, rhs).map(|(l, r)| l + &r).collect())
206			// An empty hypercube yields no partials at all, and its round evals are zero.
207			.unwrap_or_else(|| vec![RoundEvals::<P, 2>::default(); sums.len()]);
208
209		// This prover has multiple evaluation points and cannot implement MleCheckProver.
210		let (prime_coeffs, round_coeffs) = izip!(&self.eq_trackers, sums, packed_prime_evals)
211			.map(|(eq_tracker, &sum, packed_prime_evals)| {
212				eq_tracker.interpolate2(sum, packed_prime_evals.sum_scalars(self.n_vars() - 1))
213			})
214			.unzip::<_, _, Vec<_>, Vec<_>>();
215
216		self.last_coeffs_or_sums = RoundState::Coeffs(prime_coeffs);
217
218		// Combine the per-claim round polynomials into the single weighted round polynomial
219		// `Σ_i weights[i] · R_i`.
220		let combined = izip!(round_coeffs, &self.weights)
221			.map(|(coeffs, &w)| coeffs * w)
222			.sum();
223		vec![combined]
224	}
225
226	fn fold(&mut self, challenge: F) {
227		let prime_coeffs = self.last_coeffs_or_sums.coeffs();
228
229		assert!(self.n_vars() > 0);
230
231		let sums = prime_coeffs
232			.iter()
233			.map(|coeffs| coeffs.evaluate(&challenge))
234			.collect();
235
236		self.eq_trackers
237			.par_iter_mut()
238			.for_each(|eq_tracker| eq_tracker.fold(challenge));
239
240		self.switchover.fold(challenge);
241		fold_highest_var_inplace(&mut self.selected, challenge);
242
243		self.last_coeffs_or_sums = RoundState::Claim(sums);
244	}
245
246	fn finish(self) -> Vec<F> {
247		assert_eq!(self.n_vars(), 0, "finish called out of order; sumcheck rounds remain");
248
249		let mut multilinear_evals = Vec::with_capacity(self.eq_trackers.len() + 1);
250
251		for selector in self.switchover.finalize() {
252			debug_assert_eq!(selector.log_len(), 0);
253			let eval = selector.get(0);
254			multilinear_evals.push(eval);
255		}
256
257		debug_assert_eq!(self.selected.log_len(), 0);
258		multilinear_evals.push(self.selected.get(0));
259
260		multilinear_evals
261	}
262}
263
264#[cfg(test)]
265mod tests {
266	use std::iter::repeat_with;
267
268	use binius_field::FieldOps;
269	use binius_ip::sumcheck::verify;
270	use binius_math::{
271		multilinear::{eq::eq_ind, evaluate::evaluate},
272		test_utils::{Packed128b, random_scalars},
273	};
274	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
275	use itertools::Itertools;
276	use rand::prelude::*;
277
278	use super::*;
279	use crate::sumcheck::prove::prove_single;
280
281	type P = Packed128b;
282	type F = <P as FieldOps>::Scalar;
283	type StdChallenger = HasherChallenger<sha2::Sha256>;
284
285	// Prove/verify roundtrip: drive the prover through a transcript, verify with the generic
286	// sumcheck verifier, and reconstruct the reduced claim from the returned multilinear
287	// evaluations. This mirrors the verifier's selector-sumcheck check (`verify_phase_3` in
288	// `binius-verifier`'s intmul protocol), which recombines per-selector terms
289	// `(selector·(selected − 1) + 1)·eq(point_i, r)` weighted by an equality tensor.
290	#[test]
291	fn test_selector_mlecheck_prove_verify() {
292		let mut rng = StdRng::seed_from_u64(0);
293
294		let n_vars = 8;
295		let selector_count = 3;
296
297		let selector_mask = (1u16 << selector_count) - 1;
298		let bitmasks = repeat_with(|| rng.random::<u16>() & selector_mask)
299			.take(1 << n_vars)
300			.collect_vec();
301
302		let selected_scalars = random_scalars::<F>(&mut rng, 1 << n_vars);
303		let selected = FieldBuffer::<P>::from_values(&selected_scalars);
304
305		// The 1-bit selector columns, extracted from the bitmasks.
306		let selector_columns = (0..selector_count)
307			.map(|i| {
308				bitmasks
309					.iter()
310					.map(|b| if (b >> i) & 1 == 1 { F::ONE } else { F::ZERO })
311					.collect_vec()
312			})
313			.collect_vec();
314
315		// One claim per selector: the composition `selected * selector + (1 - selector)` evaluated
316		// at an independent random point.
317		let points = repeat_with(|| random_scalars::<F>(&mut rng, n_vars))
318			.take(selector_count)
319			.collect_vec();
320		let claims = izip!(&selector_columns, &points)
321			.map(|(selector_scalars, point)| {
322				let masked = izip!(&selected_scalars, selector_scalars)
323					.map(|(&selected, &selector)| selected * selector + (F::ONE - selector))
324					.collect_vec();
325				let value = evaluate(&FieldBuffer::<P>::from_values(&masked), point);
326				Claim {
327					point: point.clone(),
328					value,
329				}
330			})
331			.collect_vec();
332
333		let weights = random_scalars::<F>(&mut rng, selector_count);
334
335		// The prover reduces the per-selector claims to a single weighted sumcheck claim.
336		let claim: F = izip!(&claims, &weights).map(|(c, &w)| c.value * w).sum();
337
338		let switchover = 0;
339		let prover = SelectorMlecheckProver::new(
340			selected.clone(),
341			claims,
342			&bitmasks,
343			weights.clone(),
344			switchover,
345		);
346
347		// Run the prover through the transcript and append the final multilinear evaluations.
348		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
349		let output = prove_single(prover, &mut prover_transcript);
350		prover_transcript
351			.message()
352			.write_slice(&output.multilinear_evals);
353
354		// Verify against the generic sumcheck verifier. The composition has degree 3: a degree-2
355		// product (`selected * selector`) times the equality indicator.
356		let mut verifier_transcript = prover_transcript.into_verifier();
357		let sumcheck_output = verify(n_vars, 3, claim, &mut verifier_transcript).unwrap();
358
359		assert_eq!(
360			output.challenges, sumcheck_output.challenges,
361			"prover and verifier challenges must match"
362		);
363
364		// The prover binds variables high-to-low; `evaluate` and `eq_ind` expect low-to-high.
365		let mut reduced_point = sumcheck_output.challenges.clone();
366		reduced_point.reverse();
367
368		// `finish()` returns `[selector_0(r), .., selector_{k-1}(r), selected(r)]`.
369		let multilinear_evals: Vec<F> = verifier_transcript
370			.message()
371			.read_vec(selector_count + 1)
372			.unwrap();
373		let (selector_evals, selected_eval) = multilinear_evals.split_at(selector_count);
374		let selected_eval = selected_eval[0];
375
376		// The claimed evaluations must match direct evaluation of the multilinears at the challenge
377		// point.
378		assert_eq!(selected_eval, evaluate(&selected, &reduced_point), "selected evaluation");
379		for (i, (&selector_eval, selector_scalars)) in
380			izip!(selector_evals, &selector_columns).enumerate()
381		{
382			assert_eq!(
383				selector_eval,
384				evaluate(&FieldBuffer::<P>::from_values(selector_scalars), &reduced_point),
385				"selector {i} evaluation"
386			);
387		}
388
389		// Reconstruct the reduced sumcheck claim from the multilinear evaluations:
390		// `Σ_i weights[i] · (selected(r)·selector_i(r) + (1 − selector_i(r))) · eq(point_i, r)`.
391		let expected_eval: F = izip!(selector_evals, &points, &weights)
392			.map(|(&selector_eval, point, &weight)| {
393				let composition = selected_eval * selector_eval + (F::ONE - selector_eval);
394				weight * composition * eq_ind(point, &reduced_point)
395			})
396			.sum();
397		assert_eq!(
398			expected_eval, sumcheck_output.eval,
399			"reduced sumcheck claim must match the composition evaluated at the challenge point"
400		);
401	}
402}