Skip to main content

binius_ip_prover/sumcheck/
round_evals.rs

1// Copyright 2023-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Sampled round-polynomial evaluations, and their interpolation to monomial coefficients.
5//!
6//! A sumcheck round produces one univariate polynomial per claim.
7//! The prover samples it at a few nodes, then solves for its coefficients.
8//!
9//! ```text
10//!     accumulate    sum the composition over the halved hypercube, at each sampled node
11//!     reduce        collapse the wide accumulator, once per round
12//!     sum_scalars   sum the packed lanes into field elements
13//!     interpolate   recover the node at 0 from the round claim, then solve
14//! ```
15//!
16//! The evaluation at 0 is never sampled.
17//! Recovering it from the round claim saves the prover one node per round.
18
19use std::{
20	array, iter,
21	ops::{Add, AddAssign, Mul},
22};
23
24use binius_field::{Field, PackedField, WideMul};
25use binius_ip::sumcheck::RoundCoeffs;
26
27/// The `D` sampled evaluations of one claim's round polynomial.
28///
29/// A degree-`D` polynomial needs `D + 1` values.
30/// The round claim supplies the one at 0, so only `D` are sampled:
31///
32/// ```text
33///     D = 1    [ R(1) ]
34///     D = 2    [ R(1), R(inf) ]
35/// ```
36///
37/// Invariant: slot 0 is `R(1)`, which every recovery of `R(0)` reads.
38///
39/// `R(inf)` is the leading coefficient, at one addition per element.
40/// Any finite node past 1 would cost a full extrapolation.
41///
42/// `T` narrows as the round advances:
43///
44/// ```text
45///     accumulate    wide, unreduced
46///     reduce        packed
47///     sum_scalars   scalar
48/// ```
49#[derive(Clone, Copy, Debug)]
50pub struct RoundEvals<T, const D: usize>(pub [T; D]);
51
52impl<T: Default, const D: usize> Default for RoundEvals<T, D> {
53	fn default() -> Self {
54		Self(array::from_fn(|_| T::default()))
55	}
56}
57
58impl<T, const D: usize> RoundEvals<T, D> {
59	/// Reads one evaluator's run of accumulator slots.
60	///
61	/// # Panics
62	///
63	/// Panics unless the run holds exactly `D` slots.
64	/// A mismatch means the evaluator's reported degree disagrees with the `D` its body uses.
65	pub fn from_slots(slots: &[T]) -> Self
66	where
67		T: Copy,
68	{
69		Self(
70			<[T; D]>::try_from(slots)
71				.expect("slot run length must equal the evaluator's reported degree"),
72		)
73	}
74
75	/// Adds these evaluations into one evaluator's run of accumulator slots.
76	///
77	/// An evaluator sums into locals through its hot loop, then calls this once per chunk.
78	/// The running sums therefore stay in registers, not in the shared accumulator buffer.
79	///
80	/// # Panics
81	///
82	/// Panics unless the run holds exactly `D` slots, as in [`Self::from_slots`].
83	pub fn add_to(self, slots: &mut [T])
84	where
85		T: AddAssign,
86	{
87		assert_eq!(slots.len(), D, "slot run length must equal the evaluator's reported degree");
88		for (slot, eval) in iter::zip(slots, self.0) {
89			*slot += eval;
90		}
91	}
92
93	/// Reduces every wide slot, once the round's accumulation is complete.
94	///
95	/// Accumulation sums unreduced products.
96	/// So the modular reduction is paid once per round, not once per hypercube element.
97	pub fn reduce<P: PackedField + WideMul<Output = T>>(self) -> RoundEvals<P, D> {
98		RoundEvals(self.0.map(P::reduce))
99	}
100}
101
102impl<P: PackedField, const D: usize> RoundEvals<P, D> {
103	/// Sums the packed lanes of every slot into a scalar.
104	///
105	/// A round narrower than one packed word leaves the trailing lanes dead.
106	/// So only the first `2^n_vars` are summed.
107	pub fn sum_scalars(self, n_vars: usize) -> RoundEvals<P::Scalar, D> {
108		RoundEvals(self.0.map(|eval| eval.iter().take(1 << n_vars).sum()))
109	}
110}
111
112impl<F: Field, const D: usize> RoundEvals<F, D> {
113	/// Recovers `R(0)` from an MLE-check round claim, at any prime degree.
114	///
115	/// The verifier reduces with `claim = (1 - alpha) * R(0) + alpha * R(1)`.
116	///
117	/// Solving for `R(0)` divides by `1 - alpha`.
118	/// That is non-zero for an honest challenge.
119	fn eval_at_zero_eq(&self, claim: F, alpha: F) -> F {
120		(claim - self.0[0] * alpha) * (F::ONE - alpha).invert_or_zero()
121	}
122}
123
124impl<F: Field> RoundEvals<F, 1> {
125	/// Interpolates the prime round polynomial of an MLE-check.
126	///
127	/// # Arguments
128	///
129	/// * `claim` - This round's claim on the prime polynomial.
130	/// * `alpha` - The coordinate of the evaluation point that this round binds.
131	pub fn interpolate_eq(self, claim: F, alpha: F) -> RoundCoeffs<F> {
132		let y_0 = self.eval_at_zero_eq(claim, alpha);
133		let [y_1] = self.0;
134
135		// For P(X) = c_1 X + c_0:
136		//
137		//     P(0) =       c_0
138		//     P(1) = c_1 + c_0
139		RoundCoeffs(vec![y_0, y_1 - y_0])
140	}
141}
142
143impl<F: Field> RoundEvals<F, 2> {
144	/// Interpolates a plain sumcheck round polynomial.
145	///
146	/// # Arguments
147	///
148	/// * `claim` - This round's sum claim.
149	pub fn interpolate(self, claim: F) -> RoundCoeffs<F> {
150		// The verifier checks `claim = R(0) + R(1)`, so the prover never samples 0.
151		let y_0 = claim - self.0[0];
152		self.coeffs_from_zero(y_0)
153	}
154
155	/// Interpolates the prime round polynomial of an MLE-check.
156	///
157	/// # Arguments
158	///
159	/// * `claim` - This round's claim on the prime polynomial.
160	/// * `alpha` - The coordinate of the evaluation point that this round binds.
161	pub fn interpolate_eq(self, claim: F, alpha: F) -> RoundCoeffs<F> {
162		let y_0 = self.eval_at_zero_eq(claim, alpha);
163		self.coeffs_from_zero(y_0)
164	}
165
166	/// Solves for the monomial coefficients, given the recovered evaluation at 0.
167	fn coeffs_from_zero(self, y_0: F) -> RoundCoeffs<F> {
168		let [y_1, y_inf] = self.0;
169
170		// For P(X) = c_2 X^2 + c_1 X + c_0:
171		//
172		//     P(0)   =             c_0
173		//     P(1)   = c_2 + c_1 + c_0
174		//     P(inf) = c_2
175		RoundCoeffs(vec![y_0, y_1 - y_0 - y_inf, y_inf])
176	}
177}
178
179impl<T: AddAssign + Copy, const D: usize> AddAssign<&Self> for RoundEvals<T, D> {
180	fn add_assign(&mut self, rhs: &Self) {
181		for (slot, &eval) in iter::zip(&mut self.0, &rhs.0) {
182			*slot += eval;
183		}
184	}
185}
186
187impl<T: AddAssign + Copy, const D: usize> Add<&Self> for RoundEvals<T, D> {
188	type Output = Self;
189
190	fn add(mut self, rhs: &Self) -> Self::Output {
191		self += rhs;
192		self
193	}
194}
195
196impl<P: PackedField, const D: usize> Mul<P::Scalar> for RoundEvals<P, D> {
197	type Output = Self;
198
199	fn mul(mut self, rhs: P::Scalar) -> Self::Output {
200		for eval in &mut self.0 {
201			*eval *= rhs;
202		}
203		self
204	}
205}
206
207#[cfg(test)]
208mod tests {
209	use binius_field::FieldOps;
210	use binius_math::test_utils::Packed128b;
211	use proptest::prelude::*;
212
213	use super::*;
214
215	type P = Packed128b;
216	type F = <P as FieldOps>::Scalar;
217
218	#[test]
219	fn from_slots_reads_the_run_in_order() {
220		let slots = [F::from(1u128), F::from(2u128), F::from(3u128)];
221		// Only the run handed over is read, not the slots on either side of it.
222		let evals = RoundEvals::<F, 2>::from_slots(&slots[1..]);
223		assert_eq!(evals.0, [F::from(2u128), F::from(3u128)]);
224	}
225
226	#[test]
227	#[should_panic(expected = "slot run length must equal the evaluator's reported degree")]
228	fn from_slots_rejects_a_run_of_the_wrong_length() {
229		let slots = [F::from(1u128), F::from(2u128), F::from(3u128)];
230		let _ = RoundEvals::<F, 2>::from_slots(&slots);
231	}
232
233	#[test]
234	fn add_to_accumulates_into_the_run_alone() {
235		let mut slots = [F::from(1u128), F::from(2u128), F::from(3u128)];
236		// A two-slot run written at offset 1 must leave slot 0 untouched.
237		RoundEvals::<F, 2>([F::from(4u128), F::from(5u128)]).add_to(&mut slots[1..]);
238		assert_eq!(slots[0], F::from(1u128));
239		assert_eq!(slots[1], F::from(2u128) + F::from(4u128));
240		assert_eq!(slots[2], F::from(3u128) + F::from(5u128));
241	}
242
243	#[test]
244	#[should_panic(expected = "slot run length must equal the evaluator's reported degree")]
245	fn add_to_rejects_a_run_of_the_wrong_length() {
246		let mut slots = [F::from(1u128)];
247		RoundEvals::<F, 2>([F::from(2u128), F::from(3u128)]).add_to(&mut slots);
248	}
249
250	proptest! {
251		// Invariant: a plain sumcheck round polynomial reproduces the claim over its endpoints.
252		//
253		// This is the identity the verifier checks each round.
254		// The two sampled nodes must also come back out of the coefficients.
255		#[test]
256		fn interpolate_2_satisfies_the_sumcheck_identity(
257			y_1 in any::<u128>(),
258			y_inf in any::<u128>(),
259			claim in any::<u128>(),
260		) {
261			let [y_1, y_inf, claim] = [y_1, y_inf, claim].map(F::from);
262
263			let coeffs = RoundEvals::<F, 2>([y_1, y_inf]).interpolate(claim);
264
265			prop_assert_eq!(coeffs.0.len(), 3);
266			prop_assert_eq!(coeffs.sum_over_endpoints(), claim);
267			prop_assert_eq!(coeffs.evaluate(&F::ONE), y_1);
268			// The leading coefficient is by definition the evaluation at infinity.
269			prop_assert_eq!(coeffs.0[2], y_inf);
270		}
271
272		// Invariant: a degree-2 prime round polynomial reproduces the MLE-check claim.
273		//
274		// The verifier reduces with `claim = (1 - alpha) R(0) + alpha R(1)`.
275		// alpha = 1 makes that recovery singular, so an honest challenge excludes it.
276		#[test]
277		fn interpolate_eq_2_satisfies_the_mlecheck_identity(
278			y_1 in any::<u128>(),
279			y_inf in any::<u128>(),
280			claim in any::<u128>(),
281			alpha in any::<u128>(),
282		) {
283			let [y_1, y_inf, claim, alpha] = [y_1, y_inf, claim, alpha].map(F::from);
284			prop_assume!(alpha != F::ONE);
285
286			let coeffs = RoundEvals::<F, 2>([y_1, y_inf]).interpolate_eq(claim, alpha);
287
288			prop_assert_eq!(coeffs.0.len(), 3);
289			let reduced = (F::ONE - alpha) * coeffs.evaluate(&F::ZERO)
290				+ alpha * coeffs.evaluate(&F::ONE);
291			prop_assert_eq!(reduced, claim);
292			prop_assert_eq!(coeffs.evaluate(&F::ONE), y_1);
293			prop_assert_eq!(coeffs.0[2], y_inf);
294		}
295
296		// Invariant: the degree-1 prime polynomial satisfies the same MLE-check identity.
297		//
298		// One sampled node suffices, since the claim pins the second.
299		#[test]
300		fn interpolate_eq_1_satisfies_the_mlecheck_identity(
301			y_1 in any::<u128>(),
302			claim in any::<u128>(),
303			alpha in any::<u128>(),
304		) {
305			let [y_1, claim, alpha] = [y_1, claim, alpha].map(F::from);
306			prop_assume!(alpha != F::ONE);
307
308			let coeffs = RoundEvals::<F, 1>([y_1]).interpolate_eq(claim, alpha);
309
310			prop_assert_eq!(coeffs.0.len(), 2);
311			let reduced = (F::ONE - alpha) * coeffs.evaluate(&F::ZERO)
312				+ alpha * coeffs.evaluate(&F::ONE);
313			prop_assert_eq!(reduced, claim);
314			prop_assert_eq!(coeffs.evaluate(&F::ONE), y_1);
315		}
316
317		// Invariant: deferring the reduction changes nothing but the number of reductions.
318		//
319		// This is what lets an evaluator accumulate wide across a whole chunk.
320		//
321		//     reduce(w_a + w_b)  ==  reduce(w_a) + reduce(w_b)
322		#[test]
323		fn reduce_is_additive_over_the_wide_accumulator(
324			a in any::<u128>(),
325			b in any::<u128>(),
326			c in any::<u128>(),
327			d in any::<u128>(),
328		) {
329			let [a, b, c, d] = [a, b, c, d].map(|bits| P::broadcast(F::from(bits)));
330			let first = RoundEvals::<_, 2>([P::wide_mul(a, b), P::wide_mul(b, c)]);
331			let second = RoundEvals::<_, 2>([P::wide_mul(c, d), P::wide_mul(d, a)]);
332
333			// Accumulate both contributions into one run, exactly as an evaluator does.
334			let mut acc = RoundEvals::<<P as WideMul>::Output, 2>::default();
335			first.add_to(&mut acc.0);
336			second.add_to(&mut acc.0);
337			let deferred = acc.reduce::<P>();
338
339			let eager = first.reduce::<P>() + &second.reduce::<P>();
340
341			prop_assert_eq!(deferred.0[0], eager.0[0]);
342			prop_assert_eq!(deferred.0[1], eager.0[1]);
343		}
344
345		// Invariant: `sum_scalars` collapses only the live lanes of each slot.
346		#[test]
347		fn sum_scalars_sums_the_live_lanes(bits in any::<u128>(), n_vars in 0usize..=2) {
348			let value = F::from(bits);
349			let evals = RoundEvals::<P, 1>([P::broadcast(value)]);
350
351			let expected: F = (0..1 << n_vars).map(|_| value).sum();
352			prop_assert_eq!(evals.sum_scalars(n_vars).0[0], expected);
353		}
354	}
355}