binius_ip_prover/sumcheck/
quadratic_mle_evaluator.rs1use binius_compute::Allocator;
4use binius_field::{Field, PackedField, WideMul};
5use binius_ip::sumcheck::RoundCoeffs;
6use binius_math::{FieldSlice, FieldVec};
7
8use super::{
9 mle_store::{ColId, ColumnChunk, EvaluationChunk, MleStore, RoundContext},
10 round_evals::RoundEvals,
11 round_evaluator::{MleCheckRoundEvaluator, SharedMleCheckProver},
12};
13
14pub struct QuadraticMleEvaluator<Composition, InfinityComposition, const N: usize> {
28 cols: [ColId; N],
30 composition: Composition,
32 infinity_composition: InfinityComposition,
34}
35
36impl<Composition, InfinityComposition, const N: usize>
37 QuadraticMleEvaluator<Composition, InfinityComposition, N>
38{
39 pub fn new(
52 cols: [ColId; N],
53 composition: Composition,
54 infinity_composition: InfinityComposition,
55 ) -> Self {
56 assert!(N > 0);
58
59 Self {
60 cols,
61 composition,
62 infinity_composition,
63 }
64 }
65}
66
67pub fn quadratic_mlecheck_prover<
99 'alloc,
100 A,
101 F,
102 P,
103 Composition,
104 InfinityComposition,
105 const N: usize,
106>(
107 alloc: &'alloc A,
108 multilinears: [FieldVec<P, A>; N],
109 composition: Composition,
110 infinity_composition: InfinityComposition,
111 eval_point: Vec<F>,
112 eval_claim: F,
113) -> SharedMleCheckProver<'alloc, A, F, P, QuadraticMleEvaluator<Composition, InfinityComposition, N>>
114where
115 A: Allocator,
116 F: Field,
117 P: PackedField<Scalar = F>,
118 Composition: Fn([P; N]) -> P + Send + Sync,
119 InfinityComposition: Fn([P; N]) -> P + Send + Sync,
120{
121 let mut store = MleStore::new(eval_point.len(), alloc);
122 let cols = multilinears.map(|col| store.push_owned(col));
124 let evaluator = QuadraticMleEvaluator::new(cols, composition, infinity_composition);
125 SharedMleCheckProver::new(store, [(eval_claim, evaluator)], eval_point)
126}
127
128impl<F, P, Composition, InfinityComposition, const N: usize> MleCheckRoundEvaluator<F, P>
129 for QuadraticMleEvaluator<Composition, InfinityComposition, N>
130where
131 F: Field,
132 P: PackedField<Scalar = F>,
133 Composition: Fn([P; N]) -> P + Send + Sync,
134 InfinityComposition: Fn([P; N]) -> P + Send + Sync,
135{
136 fn degree(&self) -> usize {
137 2
139 }
140
141 fn accumulate(
142 &self,
143 chunk: &EvaluationChunk<'_, P>,
144 eq_ind: FieldSlice<'_, P>,
145 accum: &mut [<P as WideMul>::Output],
146 ) {
147 let cols: [&ColumnChunk<'_, P>; N] = self.cols.map(|id| chunk.col(id));
150
151 let los: [&[P]; N] = cols.map(|col| col.lo.as_ref());
157 let his: [&[P]; N] = cols.map(|col| col.hi.as_ref());
158
159 let mut y_1 = <P as WideMul>::Output::default();
160 let mut y_inf = <P as WideMul>::Output::default();
161 for (idx, &eq_i) in eq_ind.as_ref().iter().enumerate() {
162 let mut evals_1 = [P::default(); N];
164 let mut evals_inf = [P::default(); N];
165
166 for i in 0..N {
167 let lo_i = los[i][idx];
168 let hi_i = his[i][idx];
169
170 evals_1[i] = hi_i;
173 evals_inf[i] = lo_i + hi_i;
174 }
175
176 y_1 += P::wide_mul((self.composition)(evals_1), eq_i);
180 y_inf += P::wide_mul((self.infinity_composition)(evals_inf), eq_i);
181 }
182
183 RoundEvals([y_1, y_inf]).add_to(accum);
184 }
185
186 fn interpolate(
187 &self,
188 ctx: &RoundContext<'_, P>,
189 accum: &[P],
190 claim: F,
191 alpha: F,
192 ) -> RoundCoeffs<F> {
193 let n_vars_remaining = ctx.n_vars();
195 assert!(n_vars_remaining > 0);
196
197 RoundEvals::<P, 2>::from_slots(accum)
201 .sum_scalars(n_vars_remaining)
202 .interpolate_eq(claim, alpha)
203 }
204}
205
206#[cfg(test)]
207mod tests {
208 use std::{array, iter};
209
210 use binius_compute::GlobalAllocator;
211 use binius_field::{arch::OptimalPackedB128, field::FieldOps};
212 use binius_ip::mlecheck;
213 use binius_math::{
214 FieldBuffer,
215 multilinear::evaluate::evaluate,
216 test_utils::{random_field_buffer, random_scalars},
217 };
218 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
219 use itertools::Itertools;
220 use rand::prelude::*;
221
222 use super::*;
223 use crate::sumcheck::prove_single_mlecheck;
224
225 type StdChallenger = HasherChallenger<sha2::Sha256>;
226
227 fn prove_verify<F, P, const N: usize>(
232 composition: impl Fn([P; N]) -> P + Clone + Send + Sync,
233 infinity_composition: impl Fn([P; N]) -> P + Send + Sync,
234 ) where
235 F: Field,
236 P: PackedField<Scalar = F>,
237 {
238 let n_vars = 8;
239 let mut rng = StdRng::seed_from_u64(0);
240 let alloc = GlobalAllocator;
241
242 let multilinears: [_; N] = array::from_fn(|_| random_field_buffer::<P>(&mut rng, n_vars));
243
244 let composite_vals = (0..1 << n_vars.saturating_sub(P::LOG_WIDTH))
246 .map(|i| composition(array::from_fn(|j| multilinears[j].as_ref()[i])))
247 .collect_vec();
248 let composite_vals = FieldBuffer::new(n_vars, composite_vals);
249 let eval_point = random_scalars::<F>(&mut rng, n_vars);
250 let eval_claim = evaluate(&composite_vals, &eval_point);
251
252 let prover = quadratic_mlecheck_prover(
253 &alloc,
254 multilinears.clone(),
255 composition.clone(),
256 infinity_composition,
257 eval_point.clone(),
258 eval_claim,
259 );
260
261 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
262 let output = prove_single_mlecheck(prover, &mut prover_transcript);
263 prover_transcript
264 .message()
265 .write_slice(&output.multilinear_evals);
266
267 let mut verifier_transcript = prover_transcript.into_verifier();
268 let sumcheck_output = mlecheck::verify(
269 &eval_point,
270 2, eval_claim,
272 &mut verifier_transcript,
273 )
274 .unwrap();
275
276 let mut reduced_eval_point = sumcheck_output.challenges.clone();
279 reduced_eval_point.reverse();
280
281 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(N).unwrap();
282
283 let evals_packed: [P; N] = array::from_fn(|i| P::broadcast(multilinear_evals[i]));
285 assert_eq!(
286 composition(evals_packed).iter().next().unwrap(),
287 sumcheck_output.eval,
288 "composition of the column evaluations should equal the reduced evaluation"
289 );
290
291 for (multilinear, claimed_eval) in iter::zip(&multilinears, multilinear_evals) {
293 assert_eq!(evaluate(multilinear, &reduced_eval_point), claimed_eval);
294 }
295
296 assert_eq!(
297 output.challenges, sumcheck_output.challenges,
298 "prover and verifier challenges should match"
299 );
300 }
301
302 #[test]
303 fn test_identity_mlecheck() {
304 prove_verify::<_, OptimalPackedB128, 1>(|[a]| a, |[_a]| OptimalPackedB128::zero());
306 }
307
308 #[test]
309 fn test_linear_mlecheck() {
310 prove_verify::<_, OptimalPackedB128, 2>(
311 |[a, b]| a + b,
312 |[_a, _b]| OptimalPackedB128::zero(),
313 );
314 }
315
316 #[test]
317 fn test_bivariate_product_mlecheck() {
318 prove_verify::<_, OptimalPackedB128, 2>(|[a, b]| a * b, |[a, b]| a * b);
319 }
320
321 #[test]
322 fn test_mul_gate_mlecheck() {
323 prove_verify::<_, OptimalPackedB128, 3>(|[a, b, c]| a * b - c, |[a, b, _c]| a * b);
324 }
325
326 #[test]
327 fn test_4_variate_composition_mlecheck() {
328 prove_verify::<_, OptimalPackedB128, 4>(
329 |[a, b, c, d]| (a + b) * (c + d),
330 |[a, b, c, d]| (a + b) * (c + d),
331 );
332 }
333}