binius_ip_prover/sumcheck/
multilinear_eval.rs1use std::iter;
5
6use binius_compute::Allocator;
7use binius_field::{Field, PackedField, WideMul};
8use binius_ip::sumcheck::RoundCoeffs;
9use binius_math::{FieldSlice, FieldVec};
10
11use super::{
12 mle_store::{ColId, EvaluationChunk, MleStore, RoundContext},
13 round_evals::RoundEvals,
14 round_evaluator::{MleCheckRoundEvaluator, SharedMleCheckProver},
15};
16
17pub struct MultilinearEvalEvaluator {
26 col: ColId,
27}
28
29impl MultilinearEvalEvaluator {
30 pub const fn new(col: ColId) -> Self {
32 Self { col }
33 }
34}
35
36impl<F, P> MleCheckRoundEvaluator<F, P> for MultilinearEvalEvaluator
37where
38 F: Field,
39 P: PackedField<Scalar = F>,
40{
41 fn degree(&self) -> usize {
42 1
45 }
46
47 fn accumulate(
48 &self,
49 chunk: &EvaluationChunk<'_, P>,
50 eq_ind: FieldSlice<'_, P>,
51 accum: &mut [<P as WideMul>::Output],
52 ) {
53 let hi = chunk.col(self.col).hi.as_ref();
56 let eq_ind = eq_ind.as_ref();
57
58 assert_eq!(hi.len(), eq_ind.len());
61
62 let mut y_1 = <P as WideMul>::Output::default();
66 for (&m_i, &eq_i) in iter::zip(hi, eq_ind) {
67 y_1 += P::wide_mul(m_i, eq_i);
68 }
69 RoundEvals([y_1]).add_to(accum);
70 }
71
72 fn interpolate(
73 &self,
74 ctx: &RoundContext<'_, P>,
75 accum: &[P],
76 claim: F,
77 alpha: F,
78 ) -> RoundCoeffs<F> {
79 let n_vars_remaining = ctx.n_vars();
82 assert!(n_vars_remaining > 0);
83
84 RoundEvals::<P, 1>::from_slots(accum)
89 .sum_scalars(n_vars_remaining)
90 .interpolate_eq(claim, alpha)
91 }
92}
93
94pub fn multilinear_eval_prover<'alloc, A, F, P>(
115 alloc: &'alloc A,
116 witness: FieldVec<P, A>,
117 eval_point: &[F],
118 eval_claim: F,
119) -> SharedMleCheckProver<'alloc, A, F, P, MultilinearEvalEvaluator>
120where
121 A: Allocator,
122 F: Field,
123 P: PackedField<Scalar = F>,
124{
125 assert_eq!(
126 witness.log_len(),
127 eval_point.len(),
128 "witness must have number of variables equal to the evaluation point length"
129 );
130
131 let mut store = MleStore::new(eval_point.len(), alloc);
133 let col = store.push_owned(witness);
134 let evaluator = MultilinearEvalEvaluator::new(col);
135 SharedMleCheckProver::new(store, [(eval_claim, evaluator)], eval_point.to_vec())
136}
137
138#[cfg(test)]
139mod tests {
140 use binius_compute::GlobalAllocator;
141 use binius_field::{
142 FieldOps, Random,
143 arch::{OptimalB128, OptimalPackedB128},
144 };
145 use binius_ip::mlecheck;
146 use binius_math::{
147 multilinear::evaluate::evaluate,
148 test_utils::{random_field_buffer, random_scalars},
149 };
150 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
151 use rand::prelude::*;
152
153 use super::*;
154 use crate::sumcheck::{
155 common::MleCheckProver, prove_single_mlecheck,
156 quadratic_mle_evaluator::quadratic_mlecheck_prover,
157 };
158
159 type F = OptimalB128;
160 type P = OptimalPackedB128;
161 type StdChallenger = HasherChallenger<sha2::Sha256>;
162
163 #[test]
169 fn test_conformance_with_quadratic_mlecheck() {
170 let mut rng = StdRng::seed_from_u64(0);
171 let n_vars = 8;
172 let alloc = GlobalAllocator;
173
174 let witness = random_field_buffer::<P>(&mut rng, n_vars);
175 let eval_point = random_scalars::<F>(&mut rng, n_vars);
176 let eval_claim = evaluate(&witness, &eval_point);
177
178 let mut eval_prover =
179 multilinear_eval_prover(&alloc, witness.clone(), &eval_point, eval_claim);
180 let mut quadratic_prover = quadratic_mlecheck_prover(
181 &alloc,
182 [witness],
183 |[a]: [P; 1]| a,
184 |[_a]: [P; 1]| P::zero(),
185 eval_point,
186 eval_claim,
187 );
188
189 for _ in 0..n_vars {
190 let eval_round = eval_prover.execute();
191 let mut quadratic_round = quadratic_prover.execute();
192 assert_eq!(eval_round.len(), 1);
193 assert_eq!(quadratic_round.len(), 1);
194
195 assert_eq!(quadratic_round[0].0.pop(), Some(F::ZERO));
198 assert_eq!(eval_round[0], quadratic_round[0]);
199
200 let challenge = F::random(&mut rng);
201 eval_prover.fold(challenge);
202 quadratic_prover.fold(challenge);
203 }
204
205 assert_eq!(eval_prover.finish(), quadratic_prover.finish());
206 }
207
208 #[test]
210 fn test_prove_verify_roundtrip() {
211 let mut rng = StdRng::seed_from_u64(1);
212 let n_vars = 7;
213 let alloc = GlobalAllocator;
214
215 let witness = random_field_buffer::<P>(&mut rng, n_vars);
216 let eval_point = random_scalars::<F>(&mut rng, n_vars);
217 let eval_claim = evaluate(&witness, &eval_point);
218
219 let prover = multilinear_eval_prover(&alloc, witness.clone(), &eval_point, eval_claim);
220
221 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
222 let output = prove_single_mlecheck(prover, &mut prover_transcript);
223 prover_transcript
224 .message()
225 .write_slice(&output.multilinear_evals);
226
227 let mut verifier_transcript = prover_transcript.into_verifier();
228 let sumcheck_output =
229 mlecheck::verify::<F, _>(&eval_point, 1, eval_claim, &mut verifier_transcript).unwrap();
230 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(1).unwrap();
231
232 assert_eq!(output.challenges, sumcheck_output.challenges);
233
234 assert_eq!(multilinear_evals[0], sumcheck_output.eval);
236
237 let mut reduced_point = sumcheck_output.challenges;
238 reduced_point.reverse();
239 assert_eq!(evaluate(&witness, &reduced_point), multilinear_evals[0]);
240 }
241}