Skip to main content

binius_ip_prover/sumcheck/
eq_tracker.rs

1// Copyright 2023-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Equality-indicator state carried across the rounds of an MLE-check.
5//!
6//! An MLE-check weights its hypercube sum by the equality indicator at a fixed point.
7//! Carrying that indicator whole would cost one buffer entry per hypercube vertex.
8//!
9//! [Gruen24] section 3.2 splits it into three factors, per round:
10//!
11//! ```text
12//!     scalar      equality terms of the coordinates already bound
13//!     linear      the term in the variable this round binds
14//!     expansion   the indicator over the coordinates still untouched
15//! ```
16//!
17//! Only the expansion needs a buffer.
18//! It holds one variable fewer than the columns, since the bound one is not expanded.
19//!
20//! The other two factors are the scalar state every tracker here carries.
21//!
22//! [Gruen24]: <https://eprint.iacr.org/2024/108>
23
24use std::cmp::min;
25
26use binius_field::{Field, PackedField};
27use binius_ip::sumcheck::RoundCoeffs;
28use binius_math::{
29	field_buffer::FieldBuffer,
30	multilinear::eq::{eq_ind_partial_eval, eq_ind_truncate_low_inplace, eq_one_var},
31};
32
33use super::round_evals::RoundEvals;
34
35/// Where the rounds have reached in the point, and the product they accrued.
36#[derive(Debug, Clone)]
37struct EqPrefix<F: Field> {
38	/// The whole point, including coordinates already bound.
39	eval_point: Vec<F>,
40	/// How many coordinates are still unbound, counting the one bound next.
41	n_vars_remaining: usize,
42	/// The product of equality terms over the coordinates already bound.
43	eq_prefix_eval: F,
44}
45
46impl<F: Field> EqPrefix<F> {
47	/// Starts at the full point, with nothing bound and an empty product.
48	fn new(eval_point: &[F]) -> Self {
49		Self {
50			eval_point: eval_point.to_vec(),
51			n_vars_remaining: eval_point.len(),
52			// An empty product is one, so the first round scales by nothing.
53			eq_prefix_eval: F::ONE,
54		}
55	}
56
57	/// Returns the coordinate of the variable the next round binds.
58	///
59	/// Rounds bind from the highest coordinate down.
60	/// So this walks the point backwards.
61	fn next_coordinate(&self) -> F {
62		self.eval_point[self.n_vars_remaining - 1]
63	}
64
65	/// Binds the current variable to `challenge`, accruing its equality term.
66	fn advance(&mut self, challenge: F) {
67		// precondition
68		assert!(self.n_vars_remaining > 0);
69
70		// The bound variable's linear term becomes a constant in the product.
71		self.eq_prefix_eval *= eq_one_var(challenge, self.next_coordinate());
72		self.n_vars_remaining -= 1;
73	}
74}
75
76/// Equality-indicator state for one point, holding the expansion as one buffer.
77///
78/// A shared column store keeps this shape.
79/// Its expansion is read in the same chunks as the columns beside it.
80#[derive(Debug, Clone)]
81pub struct EqTracker<P: PackedField> {
82	/// Where the rounds have reached, and the product they accrued.
83	prefix: EqPrefix<P::Scalar>,
84	/// The indicator over the coordinates no round has reached.
85	expansion: FieldBuffer<P>,
86}
87
88impl<F: Field, P: PackedField<Scalar = F>> EqTracker<P> {
89	/// Expands the indicator for every coordinate of `eval_point` bar the highest.
90	pub fn new(eval_point: &[F]) -> Self {
91		// The round about to run keeps the highest coordinate as a linear term.
92		// It is therefore never expanded.
93		let expanded = &eval_point[..eval_point.len().saturating_sub(1)];
94		Self {
95			prefix: EqPrefix::new(eval_point),
96			expansion: eq_ind_partial_eval(expanded),
97		}
98	}
99
100	/// Returns the whole point, including coordinates already bound.
101	pub fn eval_point(&self) -> &[F] {
102		&self.prefix.eval_point
103	}
104
105	/// Returns the coordinate of the variable the next round binds.
106	pub fn next_coordinate(&self) -> F {
107		self.prefix.next_coordinate()
108	}
109
110	/// Returns the product of equality terms over the coordinates already bound.
111	pub const fn eq_prefix_eval(&self) -> F {
112		self.prefix.eq_prefix_eval
113	}
114
115	/// Returns the expansion over the coordinates no round has reached.
116	pub const fn expansion(&self) -> &FieldBuffer<P> {
117		&self.expansion
118	}
119
120	/// Returns the expansion for a caller that contracts the values itself.
121	pub const fn expansion_mut(&mut self) -> &mut FieldBuffer<P> {
122		&mut self.expansion
123	}
124
125	/// Advances one round, shrinking the expansion and accruing the equality term.
126	///
127	/// # Arguments
128	///
129	/// * `challenge` - The value this round binds the current variable to.
130	/// * `shrink` - How the expansion drops its highest variable.
131	fn advance(&mut self, challenge: F, shrink: impl FnOnce(&mut FieldBuffer<P>, usize)) {
132		// The expansion trails the columns by the coordinate this round binds.
133		debug_assert_eq!(self.expansion.log_len(), self.prefix.n_vars_remaining - 1);
134
135		// The last round finds a lone scalar, with nothing left to drop.
136		if self.expansion.log_len() > 0 {
137			let shrunk = self.expansion.log_len() - 1;
138			shrink(&mut self.expansion, shrunk);
139		}
140
141		self.prefix.advance(challenge);
142	}
143
144	/// Advances one round, contracting the expansion onto the coordinates that remain.
145	pub fn fold(&mut self, challenge: F) {
146		// Summing the two halves marginalises out the highest variable.
147		self.advance(challenge, |expansion, shrunk| {
148			eq_ind_truncate_low_inplace(expansion, shrunk);
149		});
150	}
151
152	/// Advances one round over values the caller has already contracted.
153	///
154	/// A fused read pass sums the halves during its own traversal.
155	/// Only the bookkeeping is then owed, so the stale tail is dropped.
156	pub fn truncate_one_var(&mut self, challenge: F) {
157		self.advance(challenge, |expansion, shrunk| expansion.truncate(shrunk));
158	}
159}
160
161/// Equality-indicator state for one point, holding the expansion as an outer product.
162///
163/// A prover walking the hypercube in chunks reads the two factors at different rates:
164///
165/// ```text
166///     expansion[s * chunk_len + c]  =  chunk[c] * suffix[s]
167///
168///     chunk    read per element, so it should stay in cache
169///     suffix   read once per chunk, as a single scalar
170/// ```
171///
172/// Keeping the per-element factor small is what lets it stay resident.
173#[derive(Debug, Clone)]
174pub struct ChunkedEqTracker<P: PackedField> {
175	/// Where the rounds have reached, and the product they accrued.
176	prefix: EqPrefix<P::Scalar>,
177	/// The inner factor, over the lowest coordinates, indexed within one chunk.
178	chunk: FieldBuffer<P>,
179	/// The outer factor, over the coordinates above, indexed by chunk.
180	suffix: FieldBuffer<P>,
181}
182
183impl<F: Field, P: PackedField<Scalar = F>> ChunkedEqTracker<P> {
184	/// Expands the indicator as two factors, for every coordinate bar the highest.
185	///
186	/// # Arguments
187	///
188	/// * `max_chunk_vars` - Ceiling on the inner factor's variable count.
189	/// * `eval_point` - The point the claim is weighted at.
190	pub fn new(max_chunk_vars: usize, eval_point: &[F]) -> Self {
191		// The round about to run keeps the highest coordinate as a linear term.
192		let expanded = &eval_point[..eval_point.len().saturating_sub(1)];
193
194		// A point below the ceiling puts everything in the inner factor.
195		let chunk_vars = min(max_chunk_vars, expanded.len());
196
197		// The inner factor takes the low coordinates, which vary within a chunk.
198		let (chunk_point, suffix_point) = expanded.split_at(chunk_vars);
199		Self {
200			prefix: EqPrefix::new(eval_point),
201			chunk: eq_ind_partial_eval(chunk_point),
202			suffix: eq_ind_partial_eval(suffix_point),
203		}
204	}
205
206	/// Returns the coordinate of the variable the next round binds.
207	pub fn next_coordinate(&self) -> F {
208		self.prefix.next_coordinate()
209	}
210
211	/// Returns the inner factor, indexed within one chunk.
212	pub const fn chunk(&self) -> &FieldBuffer<P> {
213		&self.chunk
214	}
215
216	/// Returns the outer factor, indexed by chunk.
217	pub const fn suffix(&self) -> &FieldBuffer<P> {
218		&self.suffix
219	}
220
221	/// Interpolates a degree-2 round polynomial from its sampled evaluations.
222	///
223	/// # Arguments
224	///
225	/// * `sum` - The claim this round's polynomial must satisfy.
226	/// * `prime_evals` - Sampled evaluations, with the equality factors removed.
227	///
228	/// # Returns
229	///
230	/// * The polynomial without its equality factors, which the next claim reduces.
231	/// * The same polynomial carrying both factors, which goes on the wire.
232	pub fn interpolate2(
233		&self,
234		sum: F,
235		prime_evals: RoundEvals<F, 2>,
236	) -> (RoundCoeffs<F>, RoundCoeffs<F>) {
237		let alpha = self.next_coordinate();
238
239		// Dropping the equality factor lowers the degree by one.
240		// So one evaluation fewer pins the polynomial down.
241		let prime_coeffs = prime_evals.interpolate_eq(sum, alpha);
242
243		// Multiplying the linear term back in restores the second factor.
244		// Scaling by the accrued product restores the first.
245		let round_coeffs = prime_coeffs.mul_by_eq(alpha) * self.prefix.eq_prefix_eval;
246
247		(prime_coeffs, round_coeffs)
248	}
249
250	/// Advances one round, contracting whichever factor holds the bound coordinate.
251	pub fn fold(&mut self, challenge: F) {
252		// Together the factors trail the columns by the coordinate this round binds.
253		debug_assert_eq!(
254			self.chunk.log_len() + self.suffix.log_len(),
255			self.prefix.n_vars_remaining - 1
256		);
257
258		// Rounds bind downwards, and the outer factor holds the higher coordinates.
259		// So it empties before the inner factor is touched at all.
260		//
261		//     suffix non-empty  ->  shrink the outer factor
262		//     suffix empty      ->  shrink the inner factor
263		//     both empty        ->  last round, nothing to drop
264		let factor = if self.suffix.log_len() > 0 {
265			Some(&mut self.suffix)
266		} else if self.chunk.log_len() > 0 {
267			Some(&mut self.chunk)
268		} else {
269			None
270		};
271
272		// Summing the two halves marginalises out that factor's highest variable.
273		if let Some(factor) = factor {
274			let truncated_log_len = factor.log_len() - 1;
275			eq_ind_truncate_low_inplace(factor, truncated_log_len);
276		}
277
278		self.prefix.advance(challenge);
279	}
280}
281
282#[cfg(test)]
283mod tests {
284	use binius_field::FieldOps;
285	use binius_math::test_utils::{Packed128b, random_scalars};
286	use rand::{SeedableRng, rngs::StdRng};
287
288	use super::*;
289
290	type P = Packed128b;
291	type F = <P as FieldOps>::Scalar;
292
293	#[test]
294	fn fold_contracts_to_the_unbound_prefix() {
295		// Invariant: the expansion trails the columns by one variable.
296		//
297		// The round's own coordinate lives in the linear term, not in the sum.
298		// So contracting once must rebuild a fresh expansion of the shorter prefix.
299		//
300		// Fixture state: 6 coordinates, bound from the highest down.
301		//
302		//     round 0:  5 expanded,  binds z_5
303		//     round 1:  4 expanded,  binds z_4
304		//     round 5:  0 expanded,  binds z_0
305		let mut rng = StdRng::seed_from_u64(0);
306		let n_vars = 6;
307		let point = random_scalars::<F>(&mut rng, n_vars);
308		let challenges = random_scalars::<F>(&mut rng, n_vars);
309
310		let mut tracker = EqTracker::<P>::new(&point);
311		// The reference product, accrued independently of the tracker.
312		let mut expected_prefix = F::ONE;
313
314		for (round, &challenge) in challenges.iter().enumerate() {
315			let unbound = n_vars - round;
316
317			// The coordinate on offer is the highest one still unbound.
318			assert_eq!(tracker.next_coordinate(), point[unbound - 1]);
319			// Only the already-bound coordinates have entered the product.
320			assert_eq!(tracker.eq_prefix_eval(), expected_prefix);
321
322			// The expansion matches one built from scratch over the lower coordinates.
323			let expected = eq_ind_partial_eval::<P>(&point[..unbound - 1]);
324			assert_eq!(tracker.expansion().log_len(), expected.log_len());
325			for i in 0..expected.len() {
326				assert_eq!(tracker.expansion().get(i), expected.get(i), "round {round}, slot {i}");
327			}
328
329			tracker.fold(challenge);
330			expected_prefix *= eq_one_var(challenge, point[unbound - 1]);
331		}
332
333		// Every coordinate is bound, so the product covers the whole point.
334		assert_eq!(tracker.eq_prefix_eval(), expected_prefix);
335	}
336
337	#[test]
338	fn truncate_matches_fold_except_in_the_values() {
339		// Invariant: a fused read pass sums the halves during its own traversal.
340		//
341		// It then owes the tracker only the bookkeeping.
342		// So the two entries agree on everything but the buffer contents.
343		//
344		//     fold      sums the halves,  so the values contract
345		//     truncate  drops the tail,   so the values stay the original front
346		//     both      same length, same accrued product
347		let mut rng = StdRng::seed_from_u64(1);
348		let n_vars = 5;
349		let point = random_scalars::<F>(&mut rng, n_vars);
350		let challenges = random_scalars::<F>(&mut rng, n_vars);
351
352		// Two trackers over one point, advanced through opposite entries.
353		let mut folded = EqTracker::<P>::new(&point);
354		let mut truncated = EqTracker::<P>::new(&point);
355		// The values the truncating path never rewrites.
356		let original = eq_ind_partial_eval::<P>(&point[..n_vars - 1]);
357
358		for (round, &challenge) in challenges.iter().enumerate() {
359			folded.fold(challenge);
360			truncated.truncate_one_var(challenge);
361
362			// Both entries run the same bookkeeping, so it cannot diverge.
363			assert_eq!(folded.eq_prefix_eval(), truncated.eq_prefix_eval(), "round {round}");
364			assert_eq!(
365				folded.expansion().log_len(),
366				truncated.expansion().log_len(),
367				"round {round}"
368			);
369
370			// Truncation leaves the front of the original values untouched.
371			for i in 0..truncated.expansion().len() {
372				assert_eq!(
373					truncated.expansion().get(i),
374					original.get(i),
375					"round {round}, slot {i}"
376				);
377			}
378		}
379	}
380
381	#[test]
382	fn chunk_and_suffix_factor_the_expansion() {
383		// Invariant: the chunked mode never materialises the full expansion.
384		//
385		// Its consumer reads one inner value and weights it by one outer scalar.
386		// So the outer product must rebuild the full expansion after every fold.
387		//
388		// Fixture state: 7 coordinates, inner factor capped at 3.
389		//
390		//     round 0:  chunk 3 vars,  suffix 3 vars
391		//     round 3:  chunk 3 vars,  suffix 0 vars
392		//     round 6:  both a lone scalar
393		//
394		//     expansion[s * chunk_len + c] == chunk[c] * suffix[s]
395		let mut rng = StdRng::seed_from_u64(2);
396		let n_vars = 7;
397		let max_chunk_vars = 3;
398		let point = random_scalars::<F>(&mut rng, n_vars);
399		let challenges = random_scalars::<F>(&mut rng, n_vars);
400
401		let mut tracker = ChunkedEqTracker::<P>::new(max_chunk_vars, &point);
402
403		for (round, &challenge) in challenges.iter().enumerate() {
404			let unbound = n_vars - round;
405			// The reference the two factors must rebuild between them.
406			let full = eq_ind_partial_eval::<P>(&point[..unbound - 1]);
407			let chunk = tracker.chunk();
408			let suffix = tracker.suffix();
409
410			// Between them the factors cover exactly the expanded coordinates.
411			assert_eq!(chunk.log_len() + suffix.log_len(), unbound - 1, "round {round}");
412
413			// The outer factor indexes chunks, the inner one indexes within a chunk.
414			for s in 0..suffix.len() {
415				for c in 0..chunk.len() {
416					assert_eq!(
417						full.get(s * chunk.len() + c),
418						chunk.get(c) * suffix.get(s),
419						"round {round}, suffix {s}, chunk {c}"
420					);
421				}
422			}
423
424			tracker.fold(challenge);
425		}
426	}
427}