Skip to main content

binius_ip_prover/sumcheck/
batch.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use binius_field::Field;
5
6use crate::{
7	channel::IPProverChannel,
8	sumcheck::{
9		common::{MleCheckProver, SumcheckProver},
10		drive::{self, MleCheckRounds, SumcheckRounds},
11	},
12};
13
14/// Prover view of the execution result of a batched sumcheck.
15#[derive(Debug, PartialEq, Eq)]
16pub struct BatchSumcheckOutput<F: Field> {
17	/// Verifier challenges for each round of the sumcheck protocol.
18	///
19	/// One challenge is generated per variable in the multivariate polynomial,
20	/// with challenges\[i\] corresponding to the i-th round of the protocol. This binding
21	/// order matches [`prove_single`](super::prove_single) and `batch_verify`; when folding
22	/// high-to-low, reverse to obtain the variable-indexed evaluation point.
23	pub challenges: Vec<F>,
24	/// Evaluation claims on non-transparent multilinears, per prover.
25	///
26	/// Each inner vector contains the evaluation values for one prover's
27	/// multilinear polynomials at the challenge point.
28	pub multilinear_evals: Vec<Vec<F>>,
29}
30
31impl<F: Field> BatchSumcheckOutput<F> {
32	/// Sends every prover's evaluation claims to the verifier, in prover order.
33	///
34	/// The per-prover grouping is prover-side bookkeeping only.
35	/// The elements reach the transcript as one flat run, which is how the verifier reads them.
36	pub fn send_evals(&self, channel: &mut impl IPProverChannel<F>) {
37		for evals in &self.multilinear_evals {
38			// Preserve per-prover ordering when emitting evaluation claims.
39			channel.send_many(evals);
40		}
41	}
42}
43
44/// Prove a batched sumcheck protocol execution, where all provers have the same number of rounds.
45///
46/// The batched sumcheck reduces a set of claims about the sums of multivariate polynomials over
47/// the boolean hypercube to their evaluation at a (shared) challenge point. This is achieved by
48/// constructing an `n_vars + 1`-variate polynomial whose coefficients in the "new variable" are the
49/// individual sum claims and evaluating it at a random point. Due to linearity of sums each claim
50/// can be proven separately with an individual [`SumcheckProver`] followed by weighted summation of
51/// the round polynomials.
52///
53/// This function performs the sumcheck protocol and returns the challenges and evaluation claims,
54/// but does not write the evaluation claims to the channel. Use [`batch_prove_and_write_evals`]
55/// if you need to write the evaluations to the channel.
56pub fn batch_prove<F, Prover>(
57	provers: Vec<Prover>,
58	channel: &mut impl IPProverChannel<F>,
59) -> BatchSumcheckOutput<F>
60where
61	F: Field,
62	Prover: SumcheckProver<F>,
63{
64	drive::batch(provers.into_iter().map(SumcheckRounds), channel)
65}
66
67/// Prove a batched sumcheck protocol and write evaluation claims to the channel.
68///
69/// This function combines [`batch_prove`] with writing the evaluation claims to the channel.
70/// It performs the batched sumcheck protocol execution and then writes all the multilinear
71/// evaluation values to the channel in order.
72///
73/// # Arguments
74///
75/// * `provers` - Vector of sumcheck provers, each handling one claim in the batch
76/// * `channel` - The channel for sending prover messages and sampling challenges
77///
78/// # Returns
79///
80/// Returns [`BatchSumcheckOutput`] containing the challenges and evaluation claims that were
81/// written to the channel.
82pub fn batch_prove_and_write_evals<F, Prover>(
83	provers: Vec<Prover>,
84	channel: &mut impl IPProverChannel<F>,
85) -> BatchSumcheckOutput<F>
86where
87	F: Field,
88	Prover: SumcheckProver<F>,
89{
90	let output = batch_prove(provers, channel);
91	output.send_evals(channel);
92	output
93}
94
95/// Prove a batched sumcheck for MLE-check provers sharing a common evaluation point.
96///
97/// This is the MLE-check analog of [`batch_prove`]: all provers are [`MleCheckProver`]s and must
98/// agree on the same evaluation point, so the batched protocol can fold every prover with the
99/// same per-round challenge and reduce all evaluation claims via a single batching coefficient.
100pub fn batch_prove_mle<F, MleCheckProver_>(
101	provers: Vec<MleCheckProver_>,
102	channel: &mut impl IPProverChannel<F>,
103) -> BatchSumcheckOutput<F>
104where
105	F: Field,
106	MleCheckProver_: MleCheckProver<F>,
107{
108	// All MLE-check provers must share the same evaluation point to batch safely. An empty batch
109	// has no point to agree on, and the driver returns without touching the channel.
110	if let Some(first_prover) = provers.first() {
111		let eval_point = first_prover.eval_point();
112		assert!(
113			provers
114				.iter()
115				.all(|prover| prover.eval_point() == eval_point),
116			"batched MLE-check provers must share the same evaluation point"
117		);
118	}
119
120	drive::batch(provers.into_iter().map(MleCheckRounds), channel)
121}
122
123#[cfg(test)]
124mod tests {
125	use binius_field::{
126		Field, PackedField,
127		arch::{OptimalB128, OptimalPackedB128},
128	};
129	use binius_ip::sumcheck::batch_verify_mle;
130	use binius_math::{
131		FieldBuffer,
132		multilinear::evaluate::evaluate,
133		test_utils::{random_field_buffer, random_scalars},
134		univariate::evaluate_univariate,
135	};
136	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
137
138	type StdChallenger = HasherChallenger<sha2::Sha256>;
139	use binius_compute::GlobalAllocator;
140	use rand::prelude::*;
141
142	use super::batch_prove_mle;
143	use crate::sumcheck::bivariate_product_mle;
144
145	fn product_eval_claim<F, P>(
146		multilinear_a: &FieldBuffer<P>,
147		multilinear_b: &FieldBuffer<P>,
148		eval_point: &[F],
149	) -> F
150	where
151		F: Field,
152		P: PackedField<Scalar = F>,
153	{
154		let n_vars = eval_point.len();
155		let product = multilinear_a
156			.as_ref()
157			.iter()
158			.zip(multilinear_b.as_ref())
159			.map(|(&l, &r)| l * r)
160			.collect::<Vec<_>>();
161		let product_buffer = FieldBuffer::new(n_vars, product);
162		evaluate(&product_buffer, eval_point)
163	}
164
165	#[test]
166	fn test_batch_prove_verify_mlecheck() {
167		type F = OptimalB128;
168		type P = OptimalPackedB128;
169
170		let n_vars = 6;
171		let mut rng = StdRng::seed_from_u64(0);
172		let alloc = GlobalAllocator;
173
174		let eval_point = random_scalars::<F>(&mut rng, n_vars);
175
176		let multilinear_a_0 = random_field_buffer::<P>(&mut rng, n_vars);
177		let multilinear_b_0 = random_field_buffer::<P>(&mut rng, n_vars);
178		let eval_claim_0 = product_eval_claim(&multilinear_a_0, &multilinear_b_0, &eval_point);
179
180		let multilinear_a_1 = random_field_buffer::<P>(&mut rng, n_vars);
181		let multilinear_b_1 = random_field_buffer::<P>(&mut rng, n_vars);
182		let eval_claim_1 = product_eval_claim(&multilinear_a_1, &multilinear_b_1, &eval_point);
183
184		let prover_0 = bivariate_product_mle::new(
185			&alloc,
186			[multilinear_a_0.clone(), multilinear_b_0.clone()],
187			eval_point.clone(),
188			eval_claim_0,
189		);
190		let prover_1 = bivariate_product_mle::new(
191			&alloc,
192			[multilinear_a_1.clone(), multilinear_b_1.clone()],
193			eval_point.clone(),
194			eval_claim_1,
195		);
196
197		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
198		let output = batch_prove_mle(vec![prover_0, prover_1], &mut prover_transcript);
199
200		let mut writer = prover_transcript.message();
201		for evals in &output.multilinear_evals {
202			writer.write_scalar_slice(evals);
203		}
204
205		let mut verifier_transcript = prover_transcript.into_verifier();
206		let sumcheck_output = batch_verify_mle(
207			&eval_point,
208			2, // degree 2 for bivariate product MLE-check
209			&[eval_claim_0, eval_claim_1],
210			&mut verifier_transcript,
211		)
212		.unwrap();
213
214		let mut reduced_eval_point = sumcheck_output.challenges.clone();
215		reduced_eval_point.reverse();
216
217		let flattened_evals: Vec<F> = output
218			.multilinear_evals
219			.iter()
220			.flat_map(|evals| evals.iter().copied())
221			.collect();
222		let evals_from_transcript: Vec<F> = verifier_transcript.message().read_vec(4).unwrap();
223		assert_eq!(
224			flattened_evals, evals_from_transcript,
225			"Multilinear evaluations should round-trip through the transcript"
226		);
227
228		let eval_a_0 = evaluate(&multilinear_a_0, &reduced_eval_point);
229		let eval_b_0 = evaluate(&multilinear_b_0, &reduced_eval_point);
230		let eval_a_1 = evaluate(&multilinear_a_1, &reduced_eval_point);
231		let eval_b_1 = evaluate(&multilinear_b_1, &reduced_eval_point);
232
233		assert_eq!(eval_a_0, output.multilinear_evals[0][0]);
234		assert_eq!(eval_b_0, output.multilinear_evals[0][1]);
235		assert_eq!(eval_a_1, output.multilinear_evals[1][0]);
236		assert_eq!(eval_b_1, output.multilinear_evals[1][1]);
237
238		let composed_evals = vec![eval_a_0 * eval_b_0, eval_a_1 * eval_b_1];
239		let expected_batched_eval =
240			evaluate_univariate(&composed_evals, &sumcheck_output.batch_coeff);
241
242		assert_eq!(
243			expected_batched_eval, sumcheck_output.eval,
244			"Batched evaluation should match reduced evaluation"
245		);
246
247		assert_eq!(
248			output.challenges, sumcheck_output.challenges,
249			"Prover and verifier challenges should match"
250		);
251	}
252}