binius_ip_prover/prodcheck/
one_pad_mle.rs1use std::{iter, mem};
17
18use binius_compute::Allocator;
19use binius_field::{Field, PackedField};
20use binius_ip::sumcheck::RoundCoeffs;
21use binius_math::{FieldVec, multilinear::eq::eq_one_var};
22
23use crate::sumcheck::{bivariate_product_mle, common::MleCheckProver};
24
25fn select<F: Field>(s: F, v: F) -> F {
30 F::ONE + (v - F::ONE) * s
31}
32
33pub struct OnePadMleCheckProver<F: Field, Inner> {
61 eval_point: Vec<F>,
63 pad_len: usize,
65 round: usize,
67 pad_eq_prefixes: Vec<F>,
70 phase: Phase<F, Inner>,
71}
72
73enum Phase<F, Inner> {
75 Real(Inner),
77 Padding {
79 children: [F; 2],
81 bound_eq: F,
84 },
85}
86
87pub fn new<'alloc, A, F, P>(
109 alloc: &'alloc A,
110 layer: FieldVec<P, A>,
111 pad_len: usize,
112 eval_point: Vec<F>,
113 claim: F,
114) -> OnePadMleCheckProver<F, impl MleCheckProver<F> + 'alloc>
115where
116 A: Allocator,
117 F: Field,
118 P: PackedField<Scalar = F>,
119{
120 assert!(layer.log_len() >= 1); let n_real_rounds = layer.log_len() - 1;
122 assert_eq!(eval_point.len(), pad_len + n_real_rounds); let pad_eq_prefixes = iter::once(F::ONE)
127 .chain(eval_point[..pad_len].iter().scan(F::ONE, |acc, &coord| {
128 *acc *= eq_one_var(F::ZERO, coord);
129 Some(*acc)
130 }))
131 .collect::<Vec<_>>();
132 let pad_eq = pad_eq_prefixes[pad_len];
133 assert!(pad_eq != F::ZERO, "a padding coordinate of the claim point equals one");
134
135 let inner_claim = F::ONE + (claim - F::ONE) * pad_eq.invert_or_zero();
138 let inner = bivariate_product_mle::new_split_half(
139 alloc,
140 layer,
141 eval_point[pad_len..].to_vec(),
142 inner_claim,
143 );
144
145 let mut prover = OnePadMleCheckProver {
146 eval_point,
147 pad_len,
148 round: 0,
149 pad_eq_prefixes,
150 phase: Phase::Real(inner),
151 };
152 prover.advance();
154 prover
155}
156
157impl<F: Field, Inner: MleCheckProver<F>> OnePadMleCheckProver<F, Inner> {
158 const fn n_real_rounds(&self) -> usize {
160 self.eval_point.len() - self.pad_len
161 }
162
163 fn advance(&mut self) {
166 if self.round != self.n_real_rounds() || !matches!(self.phase, Phase::Real(_)) {
167 return;
168 }
169 let placeholder = Phase::Padding {
171 children: [F::ONE; 2],
172 bound_eq: F::ONE,
173 };
174 let Phase::Real(inner) = mem::replace(&mut self.phase, placeholder) else {
175 unreachable!("the guard checked the phase");
176 };
177 self.phase = Phase::Padding {
178 children: inner
179 .finish()
180 .try_into()
181 .expect("the layer prover reduces two multilinears"),
182 bound_eq: F::ONE,
183 };
184 }
185}
186
187impl<F: Field, Inner: MleCheckProver<F>> MleCheckProver<F> for OnePadMleCheckProver<F, Inner> {
188 fn n_vars(&self) -> usize {
189 self.eval_point.len() - self.round
190 }
191
192 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
193 let Self {
195 eval_point,
196 pad_len,
197 round,
198 pad_eq_prefixes,
199 phase,
200 } = self;
201 let n_vars = eval_point.len() - *round;
202
203 let coeffs = match phase {
204 Phase::Real(inner) => {
205 let mut round_coeffs = inner.execute();
206 assert_eq!(round_coeffs.len(), 1, "the layer prover carries one claim");
207 let mut coeffs = round_coeffs.pop().expect("the vector holds one element");
208 let pad_eq = pad_eq_prefixes[*pad_len];
211 coeffs *= pad_eq;
212 coeffs.0[0] += F::ONE - pad_eq;
213 coeffs
214 }
215 Phase::Padding { children, bound_eq } => {
216 let unbound_eq = pad_eq_prefixes[n_vars - 1];
218 let big_e = [*bound_eq, -*bound_eq];
221 let [[a_0, a_1], [b_0, b_1]] =
222 children.map(|child| [select(big_e[0], child), (child - F::ONE) * big_e[1]]);
223 RoundCoeffs(vec![
226 F::ONE - unbound_eq + unbound_eq * a_0 * b_0,
227 unbound_eq * (a_0 * b_1 + a_1 * b_0),
228 unbound_eq * a_1 * b_1,
229 ])
230 }
231 };
232 vec![coeffs]
233 }
234
235 fn fold(&mut self, challenge: F) {
236 match &mut self.phase {
237 Phase::Real(inner) => inner.fold(challenge),
238 Phase::Padding { bound_eq, .. } => {
239 *bound_eq *= eq_one_var(F::ZERO, challenge);
240 }
241 }
242 self.round += 1;
243 self.advance();
244 }
245
246 fn finish(self) -> Vec<F> {
247 match self.phase {
248 Phase::Padding { children, bound_eq } => {
249 children.map(|child| select(bound_eq, child)).to_vec()
250 }
251 Phase::Real(_) => panic!("finish requires every variable to be bound"),
252 }
253 }
254
255 fn eval_point(&self) -> &[F] {
256 &self.eval_point[..self.n_vars()]
257 }
258}
259
260#[cfg(test)]
264mod tests {
265 use binius_compute::GlobalAllocator;
266 use binius_field::{Random, field::FieldOps};
267 use binius_math::{
268 FieldBuffer,
269 multilinear::evaluate::evaluate,
270 test_utils::{Packed128b, random_field_buffer, random_scalars},
271 };
272 use rand::prelude::*;
273
274 use super::*;
275
276 type P = Packed128b;
277 type F = <P as FieldOps>::Scalar;
278
279 fn one_pad_layer(layer: &FieldBuffer<P>, pad_len: usize) -> FieldBuffer<P> {
284 let values = (0..1 << (layer.log_len() + pad_len))
285 .map(|index| {
286 let padding = index & ((1 << pad_len) - 1);
287 if padding == 0 {
288 layer.get(index >> pad_len)
289 } else {
290 F::ONE
291 }
292 })
293 .collect::<Vec<_>>();
294 FieldBuffer::from_values(&values)
295 }
296
297 fn split_half_claim(buffer: &FieldBuffer<P>, eval_point: &[F]) -> F {
299 let (low, high) = buffer.split_half();
300 let products = (0..low.len())
301 .map(|i| low.get(i) * high.get(i))
302 .collect::<Vec<_>>();
303 evaluate(&FieldBuffer::<P>::from_values(&products), eval_point)
304 }
305
306 fn assert_matches_padded_reference(layer: FieldBuffer<P>, pad_len: usize) {
309 let mut rng = StdRng::seed_from_u64(0);
310 let alloc = GlobalAllocator;
311
312 let padded_layer = one_pad_layer(&layer, pad_len);
313 let n_vars = padded_layer.log_len() - 1;
314 let eval_point = random_scalars::<F>(&mut rng, n_vars);
315 let claim = split_half_claim(&padded_layer, &eval_point);
316
317 let mut reference =
318 bivariate_product_mle::new_split_half(&alloc, padded_layer, eval_point.clone(), claim);
319 let mut prover = new(&alloc, layer, pad_len, eval_point, claim);
320
321 for round in 0..n_vars {
322 assert_eq!(prover.n_vars(), n_vars - round);
323 assert_eq!(prover.eval_point(), reference.eval_point());
324 assert_eq!(prover.execute(), reference.execute(), "round {round}");
325
326 let challenge = F::random(&mut rng);
327 prover.fold(challenge);
328 reference.fold(challenge);
329 }
330
331 assert_eq!(prover.finish(), reference.finish());
332 }
333
334 #[test]
335 fn matches_padded_reference() {
336 let mut rng = StdRng::seed_from_u64(1);
337 for n_real_rounds in [0, 1, 3] {
338 for pad_len in [0, 1, 3] {
339 let layer = random_field_buffer::<P>(&mut rng, n_real_rounds + 1);
340 assert_matches_padded_reference(layer, pad_len);
341 }
342 }
343 }
344
345 #[test]
349 fn matches_padded_reference_with_constant_one_child() {
350 let mut rng = StdRng::seed_from_u64(2);
351 for pad_len in [1, 2, 4] {
352 let product = F::random(&mut rng);
353 let layer = FieldBuffer::<P>::from_values(&[product, F::ONE]);
354 assert_matches_padded_reference(layer, pad_len);
355 }
356 }
357}