1use std::mem;
20
21use binius_field::Field;
22use binius_ip::sumcheck::RoundCoeffs;
23use binius_math::multilinear::eq::eq_one_var;
24
25use crate::sumcheck::common::MleCheckProver;
26
27fn select<F: Field>(s: F, v: F) -> F {
32 F::ONE + (v - F::ONE) * s
33}
34
35fn mul_linear<F: Field>([p_0, p_1]: [F; 2], [q_0, q_1]: [F; 2]) -> RoundCoeffs<F> {
37 RoundCoeffs(vec![p_0 * q_0, p_0 * q_1 + p_1 * q_0, p_1 * q_1])
38}
39
40pub struct ZeroPadMleCheckProver<F: Field, Inner> {
70 eval_point: Vec<F>,
72 pad_len: usize,
74 round: usize,
76 pad_eq_prefixes: Vec<F>,
79 phase: Phase<F, Inner>,
80}
81
82enum Phase<F, Inner> {
84 Real(Inner),
86 Padding {
88 children: [F; 4],
90 bound_eq: F,
93 },
94}
95
96pub fn unpad_claims<F: Field>(pad_eq_inv: F, claims: [F; 2]) -> [F; 2] {
108 let [num, den] = claims;
109 [num * pad_eq_inv, select(pad_eq_inv, den)]
110}
111
112pub fn new<F, Inner>(
131 pad_eq_prefixes: Vec<F>,
132 eval_point: Vec<F>,
133 inner: Inner,
134) -> ZeroPadMleCheckProver<F, Inner>
135where
136 F: Field,
137 Inner: MleCheckProver<F>,
138{
139 let pad_len = pad_eq_prefixes
140 .len()
141 .checked_sub(1)
142 .expect("precondition: non-empty");
143 assert!(eval_point.len() >= pad_len); assert_ne!(pad_eq_prefixes[pad_len], F::ZERO); assert_eq!(inner.n_vars(), eval_point.len() - pad_len); let mut prover = ZeroPadMleCheckProver {
148 eval_point,
149 pad_len,
150 round: 0,
151 pad_eq_prefixes,
152 phase: Phase::Real(inner),
153 };
154 prover.advance();
156 prover
157}
158
159pub struct ConstantFraction<F> {
167 children: [F; 4],
169}
170
171impl<F: Field> ConstantFraction<F> {
172 pub const fn new(num: F, den: F) -> Self {
174 Self {
175 children: [num, F::ZERO, den, F::ONE],
176 }
177 }
178}
179
180impl<F: Field> MleCheckProver<F> for ConstantFraction<F> {
181 fn n_vars(&self) -> usize {
182 0
183 }
184
185 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
186 panic!("a constant-fraction layer has no variables to reduce")
187 }
188
189 fn fold(&mut self, _challenge: F) {
190 panic!("a constant-fraction layer has no variables to bind")
191 }
192
193 fn finish(self) -> Vec<F> {
194 self.children.to_vec()
195 }
196
197 fn eval_point(&self) -> &[F] {
198 &[]
199 }
200}
201
202impl<F: Field, Inner: MleCheckProver<F>> ZeroPadMleCheckProver<F, Inner> {
203 const fn n_real_rounds(&self) -> usize {
205 self.eval_point.len() - self.pad_len
206 }
207
208 fn advance(&mut self) {
211 if self.round != self.n_real_rounds() || !matches!(self.phase, Phase::Real(_)) {
212 return;
213 }
214 let placeholder = Phase::Padding {
216 children: [F::ONE; 4],
217 bound_eq: F::ONE,
218 };
219 let Phase::Real(inner) = mem::replace(&mut self.phase, placeholder) else {
220 unreachable!("the guard checked the phase");
221 };
222 self.phase = Phase::Padding {
223 children: inner
224 .finish()
225 .try_into()
226 .expect("the layer prover reduces four multilinears"),
227 bound_eq: F::ONE,
228 };
229 }
230}
231
232impl<F: Field, Inner: MleCheckProver<F>> MleCheckProver<F> for ZeroPadMleCheckProver<F, Inner> {
233 fn n_vars(&self) -> usize {
234 self.eval_point.len() - self.round
235 }
236
237 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
238 let Self {
240 eval_point,
241 pad_len,
242 round,
243 pad_eq_prefixes,
244 phase,
245 } = self;
246 let n_vars = eval_point.len() - *round;
247
248 match phase {
249 Phase::Real(inner) => {
250 let mut round_coeffs = inner.execute();
251 assert_eq!(round_coeffs.len(), 2, "the layer prover carries two claims");
252 let mut den = round_coeffs.pop().expect("the vector holds two elements");
253 let mut num = round_coeffs.pop().expect("the vector holds two elements");
254 let pad_eq = pad_eq_prefixes[*pad_len];
258 num *= pad_eq;
259 den *= pad_eq;
260 den.0[0] += F::ONE - pad_eq;
261 vec![num, den]
262 }
263 Phase::Padding { children, bound_eq } => {
264 let unbound_eq = pad_eq_prefixes[n_vars - 1];
266 let big_e = [*bound_eq, -*bound_eq];
268 let [num_0, num_1, den_0, den_1] = *children;
269 let [a_0, a_1] = [num_0, num_1].map(|num| [num * big_e[0], num * big_e[1]]);
272 let [b_0, b_1] =
273 [den_0, den_1].map(|den| [select(big_e[0], den), (den - F::ONE) * big_e[1]]);
274
275 let num_coeffs = (mul_linear(a_0, b_1) + &mul_linear(a_1, b_0)) * unbound_eq;
279 let mut den_coeffs = mul_linear(b_0, b_1) * unbound_eq;
280 den_coeffs.0[0] += F::ONE - unbound_eq;
281 vec![num_coeffs, den_coeffs]
282 }
283 }
284 }
285
286 fn fold(&mut self, challenge: F) {
287 match &mut self.phase {
288 Phase::Real(inner) => inner.fold(challenge),
289 Phase::Padding { bound_eq, .. } => {
290 *bound_eq *= eq_one_var(F::ZERO, challenge);
291 }
292 }
293 self.round += 1;
294 self.advance();
295 }
296
297 fn finish(self) -> Vec<F> {
298 match self.phase {
299 Phase::Padding { children, bound_eq } => {
300 let [num_0, num_1, den_0, den_1] = children;
301 vec![
302 num_0 * bound_eq,
303 num_1 * bound_eq,
304 select(bound_eq, den_0),
305 select(bound_eq, den_1),
306 ]
307 }
308 Phase::Real(_) => panic!("finish requires every variable to be bound"),
309 }
310 }
311
312 fn eval_point(&self) -> &[F] {
313 &self.eval_point[..self.n_vars()]
314 }
315}
316
317#[cfg(test)]
321mod tests {
322 use std::iter;
323
324 use binius_compute::GlobalAllocator;
325 use binius_field::{Random, arithmetic_traits::InvertOrZero, field::FieldOps};
326 use binius_math::{
327 FieldBuffer,
328 multilinear::evaluate::evaluate,
329 test_utils::{Packed128b, random_field_buffer, random_scalars},
330 };
331 use rand::prelude::*;
332
333 use super::*;
334 use crate::sumcheck::frac_add_mle;
335
336 type P = Packed128b;
337 type F = <P as FieldOps>::Scalar;
338
339 fn pad_layer(layer: &FieldBuffer<P>, pad_len: usize, fill: F) -> FieldBuffer<P> {
345 let values = (0..1 << (layer.log_len() + pad_len))
346 .map(|index| {
347 let padding = index & ((1 << pad_len) - 1);
348 if padding == 0 {
349 layer.get(index >> pad_len)
350 } else {
351 fill
352 }
353 })
354 .collect::<Vec<_>>();
355 FieldBuffer::from_values(&values)
356 }
357
358 fn split_half_claims(num: &FieldBuffer<P>, den: &FieldBuffer<P>, eval_point: &[F]) -> [F; 2] {
360 let (num_0, num_1) = num.split_half();
361 let (den_0, den_1) = den.split_half();
362 let composite = |compose: fn(F, F, F, F) -> F| {
363 let values = (0..num_0.len())
364 .map(|i| compose(num_0.get(i), num_1.get(i), den_0.get(i), den_1.get(i)))
365 .collect::<Vec<_>>();
366 evaluate(&FieldBuffer::<P>::from_values(&values), eval_point)
367 };
368 [
369 composite(|num_0, num_1, den_0, den_1| num_0 * den_1 + num_1 * den_0),
370 composite(|_, _, den_0, den_1| den_0 * den_1),
371 ]
372 }
373
374 fn assert_matches_padded_reference(
377 rng: &mut impl Rng,
378 padded_num: FieldBuffer<P>,
379 padded_den: FieldBuffer<P>,
380 eval_point: Vec<F>,
381 claims: [F; 2],
382 mut prover: impl MleCheckProver<F>,
383 ) {
384 let alloc = GlobalAllocator;
385 let n_vars = eval_point.len();
386 let mut reference =
387 frac_add_mle::new_split_half(&alloc, padded_num, padded_den, eval_point, claims);
388
389 for round in 0..n_vars {
390 assert_eq!(prover.n_vars(), n_vars - round);
391 assert_eq!(prover.eval_point(), reference.eval_point());
392 assert_eq!(prover.execute(), reference.execute(), "round {round}");
393
394 let challenge = F::random(&mut *rng);
395 prover.fold(challenge);
396 reference.fold(challenge);
397 }
398
399 assert_eq!(prover.finish(), reference.finish());
400 }
401
402 fn pad_eq_prefixes(eval_point: &[F], pad_len: usize) -> Vec<F> {
404 iter::once(F::ONE)
405 .chain(eval_point[..pad_len].iter().scan(F::ONE, |acc, &coord| {
406 *acc *= eq_one_var(F::ZERO, coord);
407 Some(*acc)
408 }))
409 .collect()
410 }
411
412 fn padded_layer(
414 rng: &mut impl Rng,
415 num: &FieldBuffer<P>,
416 den: &FieldBuffer<P>,
417 pad_len: usize,
418 ) -> (FieldBuffer<P>, FieldBuffer<P>, Vec<F>, [F; 2]) {
419 let padded_num = pad_layer(num, pad_len, F::ZERO);
421 let padded_den = pad_layer(den, pad_len, F::ONE);
422 let eval_point = random_scalars::<F>(rng, padded_num.log_len() - 1);
423 let claims = split_half_claims(&padded_num, &padded_den, &eval_point);
424 (padded_num, padded_den, eval_point, claims)
425 }
426
427 #[test]
428 fn matches_padded_reference() {
429 let mut rng = StdRng::seed_from_u64(1);
430 let alloc = GlobalAllocator;
431
432 for n_real_rounds in [0, 1, 3] {
433 for pad_len in [0, 1, 3] {
434 let num = random_field_buffer::<P>(&mut rng, n_real_rounds + 1);
435 let den = random_field_buffer::<P>(&mut rng, n_real_rounds + 1);
436 let (padded_num, padded_den, eval_point, claims) =
437 padded_layer(&mut rng, &num, &den, pad_len);
438
439 let prefixes = pad_eq_prefixes(&eval_point, pad_len);
440 let pad_eq_inv = prefixes[pad_len].invert_or_zero();
441 let inner = frac_add_mle::new_split_half(
442 &alloc,
443 num,
444 den,
445 eval_point[pad_len..].to_vec(),
446 unpad_claims(pad_eq_inv, claims),
447 );
448 let prover = new(prefixes, eval_point.clone(), inner);
449
450 assert_matches_padded_reference(
451 &mut rng, padded_num, padded_den, eval_point, claims, prover,
452 );
453 }
454 }
455 }
456
457 #[test]
461 fn constant_fraction_matches_padded_reference() {
462 let mut rng = StdRng::seed_from_u64(2);
463
464 for pad_len in [1, 2, 4] {
465 let root_num = F::random(&mut rng);
466 let root_den = F::random(&mut rng);
467 let num = FieldBuffer::<P>::from_values(&[root_num, F::ZERO]);
468 let den = FieldBuffer::<P>::from_values(&[root_den, F::ONE]);
469 let (padded_num, padded_den, eval_point, claims) =
470 padded_layer(&mut rng, &num, &den, pad_len);
471
472 let prefixes = pad_eq_prefixes(&eval_point, pad_len);
475 let pad_eq_inv = prefixes[pad_len].invert_or_zero();
476 assert_eq!(unpad_claims(pad_eq_inv, claims), [root_num, root_den]);
477 let prover =
478 new(prefixes, eval_point.clone(), ConstantFraction::new(root_num, root_den));
479
480 assert_matches_padded_reference(
481 &mut rng, padded_num, padded_den, eval_point, claims, prover,
482 );
483 }
484 }
485}