1use binius_compute::BufferData;
5use binius_field::{Field, PackedField, WideMul};
6use binius_ip::sumcheck::RoundCoeffs;
7use binius_math::{FieldBuffer, multilinear::fold::fold_highest_var_inplace};
8use binius_utils::{bitwise::Bitwise, rayon::prelude::*};
9use itertools::izip;
10
11use super::{
12 common::SumcheckProver, eq_tracker::ChunkedEqTracker, round_evals::RoundEvals,
13 round_state::RoundState, switchover::BinarySwitchover,
14};
15
16pub struct Claim<F: Field> {
17 pub point: Vec<F>,
18 pub value: F,
19}
20
21pub struct SelectorMlecheckProver<'b, P: PackedField, B: Bitwise, Data: BufferData<P> = Vec<P>> {
35 last_coeffs_or_sums: RoundState<Vec<RoundCoeffs<P::Scalar>>, Vec<P::Scalar>>,
36 selected: FieldBuffer<P, Data>,
37 eq_trackers: Vec<ChunkedEqTracker<P>>,
38 weights: Vec<P::Scalar>,
39 switchover: BinarySwitchover<'b, P, B>,
40}
41
42impl<'b, F: Field, P: PackedField<Scalar = F>, B: Bitwise, Data: BufferData<P>>
43 SelectorMlecheckProver<'b, P, B, Data>
44{
45 pub fn new(
54 selected: FieldBuffer<P, Data>,
55 claims: Vec<Claim<F>>,
56 bitmasks: &'b [B],
57 weights: Vec<F>,
58 switchover: usize,
59 ) -> Self {
60 let n_vars = selected.log_len();
61
62 assert!(
63 claims.iter().all(|claim| claim.point.len() == n_vars),
64 "multilinears must have equal number of variables"
65 );
66
67 assert_eq!(
68 weights.len(),
69 claims.len(),
70 "number of weights must match the number of claims"
71 );
72
73 assert_eq!(
74 bitmasks.len(),
75 selected.len(),
76 "bitmasks slice length must match the selected multilinear length"
77 );
78
79 const MAX_CHUNK_VARS: usize = 8;
80 let (eq_trackers, sums) = claims
81 .into_par_iter()
82 .map(|Claim { point, value }| (ChunkedEqTracker::new(MAX_CHUNK_VARS, &point), value))
83 .collect::<(Vec<_>, Vec<_>)>();
84
85 let switchover = BinarySwitchover::new(sums.len(), switchover.min(n_vars), bitmasks);
86 let last_coeffs_or_sums = RoundState::Claim(sums);
87
88 Self {
89 last_coeffs_or_sums,
90 selected,
91 eq_trackers,
92 weights,
93 switchover,
94 }
95 }
96}
97
98impl<'b, F, P, B, Data> SumcheckProver<F> for SelectorMlecheckProver<'b, P, B, Data>
99where
100 F: Field,
101 P: PackedField<Scalar = F>,
102 B: Bitwise,
103 Data: BufferData<P>,
104{
105 fn n_vars(&self) -> usize {
106 self.selected.log_len()
107 }
108
109 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
110 let sums = self.last_coeffs_or_sums.claim();
111
112 assert!(self.n_vars() > 0);
113
114 let chunk_vars = self
124 .eq_trackers
125 .first()
126 .map(|eq_tracker| eq_tracker.chunk().log_len())
127 .unwrap_or_default();
128 let chunk_count = 1 << (self.n_vars() - 1 - chunk_vars);
129
130 let (selected_0, selected_1) = self.selected.split_half();
133
134 let packed_prime_evals = (0..chunk_count)
135 .into_par_iter()
136 .fold(
137 || {
138 (
139 vec![RoundEvals::<P, 2>::default(); sums.len()],
140 FieldBuffer::<P>::zeros(chunk_vars),
141 FieldBuffer::<P>::zeros(chunk_vars),
142 )
143 },
144 |(mut packed_prime_evals, mut binary_chunk_0, mut binary_chunk_1), chunk_index| {
145 let selected_0_chunk = selected_0.chunk(chunk_vars, chunk_index);
146 let selected_1_chunk = selected_1.chunk(chunk_vars, chunk_index);
147
148 for (bit_offset, (round_evals, eq_tracker)) in
149 izip!(&mut packed_prime_evals, &self.eq_trackers).enumerate()
150 {
151 let eq_chunk = eq_tracker.chunk();
152 let eq_suffix_eval = eq_tracker.suffix().get(chunk_index);
153
154 let selector_0_chunk = self.switchover.get_chunk(
155 &mut binary_chunk_0,
156 bit_offset,
157 chunk_vars,
158 chunk_index,
159 );
160
161 let selector_1_chunk = self.switchover.get_chunk(
162 &mut binary_chunk_1,
163 bit_offset,
164 chunk_vars,
165 chunk_index | chunk_count,
166 );
167
168 let mut wide_y_1 = <P as WideMul>::Output::default();
173 let mut wide_y_inf = <P as WideMul>::Output::default();
174 for (&eq_i, &selected_0_i, &selected_1_i, &selector_0_i, &selector_1_i) in izip!(
175 eq_chunk.as_ref(),
176 selected_0_chunk.as_ref(),
177 selected_1_chunk.as_ref(),
178 selector_0_chunk.as_ref(),
179 selector_1_chunk.as_ref(),
180 ) {
181 let selected_inf_i = selected_0_i + selected_1_i;
182 let selector_inf_i = selector_0_i + selector_1_i;
183
184 let y_1_prod = selector_1_i * (selected_1_i - P::one()) + P::one();
188 let y_inf_prod = selector_inf_i * selected_inf_i;
189 wide_y_1 += P::wide_mul(eq_i, y_1_prod);
190 wide_y_inf += P::wide_mul(eq_i, y_inf_prod);
191 }
192 let chunk_round_evals = RoundEvals([wide_y_1, wide_y_inf]).reduce::<P>();
193
194 *round_evals += &(chunk_round_evals * eq_suffix_eval);
197 }
198
199 (packed_prime_evals, binary_chunk_0, binary_chunk_1)
200 },
201 )
202 .map(|(evals, _, _)| evals)
203 .reduce_with(|lhs, rhs| izip!(lhs, rhs).map(|(l, r)| l + &r).collect())
206 .unwrap_or_else(|| vec![RoundEvals::<P, 2>::default(); sums.len()]);
208
209 let (prime_coeffs, round_coeffs) = izip!(&self.eq_trackers, sums, packed_prime_evals)
211 .map(|(eq_tracker, &sum, packed_prime_evals)| {
212 eq_tracker.interpolate2(sum, packed_prime_evals.sum_scalars(self.n_vars() - 1))
213 })
214 .unzip::<_, _, Vec<_>, Vec<_>>();
215
216 self.last_coeffs_or_sums = RoundState::Coeffs(prime_coeffs);
217
218 let combined = izip!(round_coeffs, &self.weights)
221 .map(|(coeffs, &w)| coeffs * w)
222 .sum();
223 vec![combined]
224 }
225
226 fn fold(&mut self, challenge: F) {
227 let prime_coeffs = self.last_coeffs_or_sums.coeffs();
228
229 assert!(self.n_vars() > 0);
230
231 let sums = prime_coeffs
232 .iter()
233 .map(|coeffs| coeffs.evaluate(&challenge))
234 .collect();
235
236 self.eq_trackers
237 .par_iter_mut()
238 .for_each(|eq_tracker| eq_tracker.fold(challenge));
239
240 self.switchover.fold(challenge);
241 fold_highest_var_inplace(&mut self.selected, challenge);
242
243 self.last_coeffs_or_sums = RoundState::Claim(sums);
244 }
245
246 fn finish(self) -> Vec<F> {
247 assert_eq!(self.n_vars(), 0, "finish called out of order; sumcheck rounds remain");
248
249 let mut multilinear_evals = Vec::with_capacity(self.eq_trackers.len() + 1);
250
251 for selector in self.switchover.finalize() {
252 debug_assert_eq!(selector.log_len(), 0);
253 let eval = selector.get(0);
254 multilinear_evals.push(eval);
255 }
256
257 debug_assert_eq!(self.selected.log_len(), 0);
258 multilinear_evals.push(self.selected.get(0));
259
260 multilinear_evals
261 }
262}
263
264#[cfg(test)]
265mod tests {
266 use std::iter::repeat_with;
267
268 use binius_field::FieldOps;
269 use binius_ip::sumcheck::verify;
270 use binius_math::{
271 multilinear::{eq::eq_ind, evaluate::evaluate},
272 test_utils::{Packed128b, random_scalars},
273 };
274 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
275 use itertools::Itertools;
276 use rand::prelude::*;
277
278 use super::*;
279 use crate::sumcheck::prove::prove_single;
280
281 type P = Packed128b;
282 type F = <P as FieldOps>::Scalar;
283 type StdChallenger = HasherChallenger<sha2::Sha256>;
284
285 #[test]
291 fn test_selector_mlecheck_prove_verify() {
292 let mut rng = StdRng::seed_from_u64(0);
293
294 let n_vars = 8;
295 let selector_count = 3;
296
297 let selector_mask = (1u16 << selector_count) - 1;
298 let bitmasks = repeat_with(|| rng.random::<u16>() & selector_mask)
299 .take(1 << n_vars)
300 .collect_vec();
301
302 let selected_scalars = random_scalars::<F>(&mut rng, 1 << n_vars);
303 let selected = FieldBuffer::<P>::from_values(&selected_scalars);
304
305 let selector_columns = (0..selector_count)
307 .map(|i| {
308 bitmasks
309 .iter()
310 .map(|b| if (b >> i) & 1 == 1 { F::ONE } else { F::ZERO })
311 .collect_vec()
312 })
313 .collect_vec();
314
315 let points = repeat_with(|| random_scalars::<F>(&mut rng, n_vars))
318 .take(selector_count)
319 .collect_vec();
320 let claims = izip!(&selector_columns, &points)
321 .map(|(selector_scalars, point)| {
322 let masked = izip!(&selected_scalars, selector_scalars)
323 .map(|(&selected, &selector)| selected * selector + (F::ONE - selector))
324 .collect_vec();
325 let value = evaluate(&FieldBuffer::<P>::from_values(&masked), point);
326 Claim {
327 point: point.clone(),
328 value,
329 }
330 })
331 .collect_vec();
332
333 let weights = random_scalars::<F>(&mut rng, selector_count);
334
335 let claim: F = izip!(&claims, &weights).map(|(c, &w)| c.value * w).sum();
337
338 let switchover = 0;
339 let prover = SelectorMlecheckProver::new(
340 selected.clone(),
341 claims,
342 &bitmasks,
343 weights.clone(),
344 switchover,
345 );
346
347 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
349 let output = prove_single(prover, &mut prover_transcript);
350 prover_transcript
351 .message()
352 .write_slice(&output.multilinear_evals);
353
354 let mut verifier_transcript = prover_transcript.into_verifier();
357 let sumcheck_output = verify(n_vars, 3, claim, &mut verifier_transcript).unwrap();
358
359 assert_eq!(
360 output.challenges, sumcheck_output.challenges,
361 "prover and verifier challenges must match"
362 );
363
364 let mut reduced_point = sumcheck_output.challenges.clone();
366 reduced_point.reverse();
367
368 let multilinear_evals: Vec<F> = verifier_transcript
370 .message()
371 .read_vec(selector_count + 1)
372 .unwrap();
373 let (selector_evals, selected_eval) = multilinear_evals.split_at(selector_count);
374 let selected_eval = selected_eval[0];
375
376 assert_eq!(selected_eval, evaluate(&selected, &reduced_point), "selected evaluation");
379 for (i, (&selector_eval, selector_scalars)) in
380 izip!(selector_evals, &selector_columns).enumerate()
381 {
382 assert_eq!(
383 selector_eval,
384 evaluate(&FieldBuffer::<P>::from_values(selector_scalars), &reduced_point),
385 "selector {i} evaluation"
386 );
387 }
388
389 let expected_eval: F = izip!(selector_evals, &points, &weights)
392 .map(|(&selector_eval, point, &weight)| {
393 let composition = selected_eval * selector_eval + (F::ONE - selector_eval);
394 weight * composition * eq_ind(point, &reduced_point)
395 })
396 .sum();
397 assert_eq!(
398 expected_eval, sumcheck_output.eval,
399 "reduced sumcheck claim must match the composition evaluated at the challenge point"
400 );
401 }
402}