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}