binius_ip_prover/sumcheck/
padded.rs1use binius_field::Field;
4use binius_ip::sumcheck::RoundCoeffs;
5use binius_math::multilinear::eq::eq_one_var;
6
7use crate::sumcheck::common::SumcheckProver;
8
9#[derive(Debug, Clone)]
37pub struct PaddedSumcheckDecorator<F: Field, Inner> {
38 inner: Inner,
39 n_extra_vars: usize,
40 degree: usize,
42 round: usize,
44 eq_prefix: F,
46 claims: Vec<F>,
51}
52
53impl<F: Field, Inner: SumcheckProver<F>> PaddedSumcheckDecorator<F, Inner> {
54 pub const fn new(inner: Inner, n_extra_vars: usize, claims: Vec<F>, degree: usize) -> Self {
64 Self {
65 inner,
66 n_extra_vars,
67 degree,
68 round: 0,
69 eq_prefix: F::ONE,
70 claims,
71 }
72 }
73
74 const fn in_padding_phase(&self) -> bool {
76 self.round < self.n_extra_vars
77 }
78}
79
80impl<F: Field, Inner: SumcheckProver<F>> SumcheckProver<F> for PaddedSumcheckDecorator<F, Inner> {
81 fn n_vars(&self) -> usize {
82 self.inner.n_vars() + self.n_extra_vars.saturating_sub(self.round)
83 }
84
85 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
86 if self.in_padding_phase() {
87 self.claims
90 .iter()
91 .map(|&claim| {
92 let scaled = claim * self.eq_prefix;
93 let mut coeffs = vec![F::ZERO; self.degree + 1];
94 coeffs[0] = scaled;
95 coeffs[1] = -scaled;
96 RoundCoeffs(coeffs)
97 })
98 .collect()
99 } else {
100 self.inner
101 .execute()
102 .into_iter()
103 .map(|coeffs| coeffs * self.eq_prefix)
104 .collect()
105 }
106 }
107
108 fn fold(&mut self, challenge: F) {
109 if self.in_padding_phase() {
110 self.eq_prefix *= eq_one_var(F::ZERO, challenge);
112 } else {
113 self.inner.fold(challenge);
114 }
115 self.round += 1;
116 }
117
118 fn finish(self) -> Vec<F> {
119 self.inner.finish()
122 }
123}
124
125#[cfg(test)]
126mod tests {
127 use binius_compute::GlobalAllocator;
128 use binius_field::{
129 Random,
130 arch::{OptimalB128, OptimalPackedB128},
131 };
132 use binius_ip::sumcheck::{BatchSumcheckOutput, RoundCoeffs, batch_verify, verify};
133 use binius_math::{
134 inner_product::inner_product_par,
135 multilinear::{eq::eq_one_var, evaluate::evaluate},
136 test_utils::random_field_buffer,
137 };
138 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
139 use either::Either;
140 use rand::prelude::*;
141
142 use super::*;
143 use crate::sumcheck::{
144 batch::batch_prove_and_write_evals,
145 bivariate_product_evaluator::{BivariateProductEvaluator, bivariate_product_prover},
146 prove::prove_single,
147 round_evaluator::SharedSumcheckProver,
148 };
149
150 type F = OptimalB128;
151 type P = OptimalPackedB128;
152 type StdChallenger = HasherChallenger<sha2::Sha256>;
153
154 fn make_inner<'alloc>(
155 rng: &mut impl Rng,
156 alloc: &'alloc GlobalAllocator,
157 n_vars: usize,
158 ) -> (SharedSumcheckProver<'alloc, GlobalAllocator, P, BivariateProductEvaluator>, F) {
159 let a = random_field_buffer::<P>(&mut *rng, n_vars);
160 let b = random_field_buffer::<P>(&mut *rng, n_vars);
161 let sum = inner_product_par(&a, &b);
162 let prover = bivariate_product_prover(alloc, [a, b], sum);
163 (prover, sum)
164 }
165
166 struct MultilinearSumProver(Vec<F>);
168
169 impl SumcheckProver<F> for MultilinearSumProver {
170 fn n_vars(&self) -> usize {
171 self.0.len().ilog2() as usize
172 }
173
174 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
175 let (lo, hi) = self.0.split_at(self.0.len() / 2);
177 let r_0 = lo.iter().copied().sum::<F>();
178 let r_1 = hi.iter().copied().sum::<F>();
179 vec![RoundCoeffs(vec![r_0, r_1 - r_0])]
180 }
181
182 fn fold(&mut self, challenge: F) {
183 let half = self.0.len() / 2;
184 let (lo, hi) = self.0.split_at_mut(half);
185 for (lo_i, &hi_i) in lo.iter_mut().zip(hi.iter()) {
186 *lo_i += challenge * (hi_i - *lo_i);
187 }
188 self.0.truncate(half);
189 }
190
191 fn finish(self) -> Vec<F> {
192 self.0
193 }
194 }
195
196 #[test]
199 fn test_round_polynomials_closed_form() {
200 let mut rng = StdRng::seed_from_u64(0);
201 let n_vars = 6;
202 let n_extra_vars = 3;
203 let alloc = GlobalAllocator;
204
205 let (inner, sum) = make_inner(&mut rng, &alloc, n_vars);
206 let (mut bare_inner, _) = make_inner(&mut StdRng::seed_from_u64(0), &alloc, n_vars);
209 let mut padded = PaddedSumcheckDecorator::new(inner, n_extra_vars, vec![sum], 2);
210
211 let challenges = (0..n_vars + n_extra_vars)
212 .map(|_| F::random(&mut rng))
213 .collect::<Vec<_>>();
214
215 let mut eq_prefix = F::ONE;
216 for (i, &challenge) in challenges.iter().enumerate() {
217 assert_eq!(padded.n_vars(), n_vars + n_extra_vars - i);
218
219 let round_coeffs = padded.execute();
220 assert_eq!(round_coeffs.len(), 1);
221
222 if i < n_extra_vars {
223 let v = sum * eq_prefix;
225 assert_eq!(round_coeffs[0], RoundCoeffs(vec![v, -v, F::ZERO]));
226 } else {
227 let inner_coeffs = bare_inner.execute();
229 let expected = inner_coeffs[0].clone() * eq_prefix;
230 assert_eq!(round_coeffs[0], expected);
231 bare_inner.fold(challenge);
232 }
233
234 padded.fold(challenge);
235 if i < n_extra_vars {
236 eq_prefix *= eq_one_var(F::ZERO, challenge);
237 }
238 }
239
240 assert_eq!(padded.n_vars(), 0);
241 }
242
243 #[test]
244 fn test_round_polynomials_preserve_the_claim() {
245 let mut rng = StdRng::seed_from_u64(1);
250 let n_vars = 5;
251 let n_extra_vars = 2;
252 let alloc = GlobalAllocator;
253
254 let (inner, sum) = make_inner(&mut rng, &alloc, n_vars);
255 let mut padded = PaddedSumcheckDecorator::new(inner, n_extra_vars, vec![sum], 2);
256
257 let mut running = sum;
259 for _ in 0..n_vars + n_extra_vars {
260 let round_coeffs = padded.execute();
261 assert_eq!(round_coeffs.len(), 1);
262 assert_eq!(running, round_coeffs[0].sum_over_endpoints());
263
264 let challenge = F::random(&mut rng);
266 running = round_coeffs[0].evaluate(&challenge);
267 padded.fold(challenge);
268 }
269 }
270
271 #[test]
273 fn test_prove_verify_roundtrip() {
274 let mut rng = StdRng::seed_from_u64(2);
275 let n_vars = 7;
276 let n_extra_vars = 4;
277 let total_vars = n_vars + n_extra_vars;
278 let alloc = GlobalAllocator;
279
280 let a = random_field_buffer::<P>(&mut rng, n_vars);
281 let b = random_field_buffer::<P>(&mut rng, n_vars);
282 let sum = inner_product_par(&a, &b);
283 let inner = bivariate_product_prover(&alloc, [a.clone(), b.clone()], sum);
284 let padded = PaddedSumcheckDecorator::new(inner, n_extra_vars, vec![sum], 2);
285
286 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
287 let output = prove_single(padded, &mut prover_transcript);
288 prover_transcript
289 .message()
290 .write_slice(&output.multilinear_evals);
291
292 let mut verifier_transcript = prover_transcript.into_verifier();
293 let sumcheck_output = verify(total_vars, 2, sum, &mut verifier_transcript)
294 .expect("verification should succeed");
295 let challenges = sumcheck_output.challenges;
296
297 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
298 assert_eq!(output.multilinear_evals, multilinear_evals);
299 assert_eq!(output.challenges, challenges);
300
301 let eq_pad: F = challenges[..n_extra_vars]
303 .iter()
304 .map(|&r| eq_one_var(F::ZERO, r))
305 .product();
306
307 assert_eq!(multilinear_evals[0] * multilinear_evals[1] * eq_pad, sumcheck_output.eval);
309
310 let mut inner_point = challenges[n_extra_vars..].to_vec();
312 inner_point.reverse();
313 assert_eq!(evaluate(&a, &inner_point), multilinear_evals[0]);
314 assert_eq!(evaluate(&b, &inner_point), multilinear_evals[1]);
315 }
316
317 #[test]
321 fn test_batch_padded_with_longer_lower_degree_prover() {
322 let mut rng = StdRng::seed_from_u64(4);
323 let n_vars = 4;
324 let n_extra_vars = 3;
325 let total_vars = n_vars + n_extra_vars;
326 let alloc = GlobalAllocator;
327
328 let (inner, product_sum) = make_inner(&mut rng, &alloc, n_vars);
329 let padded = PaddedSumcheckDecorator::new(inner, n_extra_vars, vec![product_sum], 2);
330
331 let m = (0..1 << total_vars)
332 .map(|_| F::random(&mut rng))
333 .collect::<Vec<_>>();
334 let m_sum = m.iter().copied().sum::<F>();
335 let linear = MultilinearSumProver(m);
336
337 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
338 let output = batch_prove_and_write_evals(
339 vec![Either::Left(padded), Either::Right(linear)],
340 &mut prover_transcript,
341 );
342
343 let mut verifier_transcript = prover_transcript.into_verifier();
344 let BatchSumcheckOutput {
345 batch_coeff,
346 eval,
347 challenges,
348 } = batch_verify(total_vars, 2, &[product_sum, m_sum], &mut verifier_transcript)
349 .expect("verification should succeed");
350 let evals: Vec<F> = verifier_transcript.message().read_vec(3).unwrap();
351 assert_eq!(output.challenges, challenges);
352
353 let eq_pad: F = challenges[..n_extra_vars]
354 .iter()
355 .map(|&r| eq_one_var(F::ZERO, r))
356 .product();
357
358 assert_eq!(evals[0] * evals[1] * eq_pad + batch_coeff * evals[2], eval);
360 }
361
362 #[test]
364 fn test_no_padding_passthrough() {
365 let mut rng = StdRng::seed_from_u64(3);
366 let n_vars = 6;
367 let alloc = GlobalAllocator;
368
369 let (inner, sum) = make_inner(&mut rng, &alloc, n_vars);
370 let padded = PaddedSumcheckDecorator::new(inner, 0, vec![sum], 2);
371 assert_eq!(padded.n_vars(), n_vars);
372
373 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
374 let output = prove_single(padded, &mut prover_transcript);
375 prover_transcript
376 .message()
377 .write_slice(&output.multilinear_evals);
378
379 let mut verifier_transcript = prover_transcript.into_verifier();
380 let sumcheck_output =
381 verify(n_vars, 2, sum, &mut verifier_transcript).expect("verification should succeed");
382 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
383
384 assert_eq!(multilinear_evals[0] * multilinear_evals[1], sumcheck_output.eval);
385 assert_eq!(output.challenges, sumcheck_output.challenges);
386 }
387}