binius_ip_prover/sumcheck/
batch.rs1use binius_field::Field;
5
6use crate::{
7 channel::IPProverChannel,
8 sumcheck::{
9 common::{MleCheckProver, SumcheckProver},
10 drive::{self, MleCheckRounds, SumcheckRounds},
11 },
12};
13
14#[derive(Debug, PartialEq, Eq)]
16pub struct BatchSumcheckOutput<F: Field> {
17 pub challenges: Vec<F>,
24 pub multilinear_evals: Vec<Vec<F>>,
29}
30
31impl<F: Field> BatchSumcheckOutput<F> {
32 pub fn send_evals(&self, channel: &mut impl IPProverChannel<F>) {
37 for evals in &self.multilinear_evals {
38 channel.send_many(evals);
40 }
41 }
42}
43
44pub 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
67pub 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
95pub 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 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, &[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}