binius_ip_prover/sumcheck/
round_evals.rs1use std::{
20 array, iter,
21 ops::{Add, AddAssign, Mul},
22};
23
24use binius_field::{Field, PackedField, WideMul};
25use binius_ip::sumcheck::RoundCoeffs;
26
27#[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 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 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 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 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 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 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 RoundCoeffs(vec![y_0, y_1 - y_0])
140 }
141}
142
143impl<F: Field> RoundEvals<F, 2> {
144 pub fn interpolate(self, claim: F) -> RoundCoeffs<F> {
150 let y_0 = claim - self.0[0];
152 self.coeffs_from_zero(y_0)
153 }
154
155 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 fn coeffs_from_zero(self, y_0: F) -> RoundCoeffs<F> {
168 let [y_1, y_inf] = self.0;
169
170 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 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 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 #[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 prop_assert_eq!(coeffs.0[2], y_inf);
270 }
271
272 #[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 #[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 #[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 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 #[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}