Skip to main content

binius_ip_prover/fracaddcheck/
padding.rs

1// Copyright 2026 The Binius Developers
2
3//! Zero-fraction padding for batching fractional-addition trees of unequal depths.
4//!
5//! A batch runs one uniform layer schedule, so every tree in it must be of the same depth.
6//! Padding lifts a tree of depth $m$ to the batch's depth $n \ge m$.
7//! The $n - m$ extra leaf positions hold the zero fraction $0/1$, the additive identity.
8//! So the tree's fractional sum is unchanged and the verifier never learns the individual depths.
9//!
10//! The padding variables are the lowest ones.
11//! Padding a witness $(N, D)$ over $\nu = n - m$ of them gives
12//!
13//! $$
14//! N'(X_\text{pad}, X_\text{real}) = N(X_\text{real}) \cdot \text{eq}(0^\nu; X_\text{pad}),
15//! \qquad
16//! D'(X_\text{pad}, X_\text{real}) = 1 + \bigl( D(X_\text{real}) - 1 \bigr) \cdot
17//! \text{eq}(0^\nu; X_\text{pad}),
18//! $$
19//!
20//! so the numerators are zero-padded and the denominators one-padded.
21//!
22//! The prover never materializes a padded witness.
23//! `PaddedBatch` holds the trees and how deep each one sits.
24//! Every layer it pops hands out one `PaddedLayerProver` per tree.
25//!
26//! Each of those wraps the tree's own layer prover in a [`ZeroPadMleCheckProver`].
27//! That wrapper corrects the unpadded layer's messages at a cost of $O(1)$ per round.
28//!
29//! A tree the batch has not reached yet has no layer to wrap.
30//! It contributes a [`ConstantFraction`] instead.
31//!
32//! [`unpad_leaf_claim`] inverts the identity above on the claims the batch outputs.
33
34use std::iter;
35
36use binius_compute::Allocator;
37use binius_field::{Field, PackedField};
38use binius_ip::{fracaddcheck::FracAddEvalClaim, sumcheck::RoundCoeffs};
39use binius_math::{batch_invert::BatchInversion, multilinear::eq::eq_one_var};
40use either::Either;
41use itertools::izip;
42
43use super::{
44	FracAddCircuit, LayerProver,
45	fraction::Fraction,
46	zero_pad_mle::{self, ConstantFraction, ZeroPadMleCheckProver},
47};
48use crate::sumcheck::common::MleCheckProver;
49
50/// The layer one tree contributes to the batch's current depth.
51///
52/// Once the batch has reached the tree, that is a real layer of it.
53///
54/// Until then it is the layer standing in for one.
55/// That stand-in is the tree's own fractional sum beside the zero fraction $0/1$.
56type TreeLayer<'a, A, F, P> = Either<LayerProver<'a, A, F, P>, ConstantFraction<F>>;
57
58/// One tree's contribution to a batched layer, lifted to the batch's depth.
59///
60/// Either kind of [`TreeLayer`] is a layer of the padded tree.
61///
62/// So either one needs its messages corrected.
63/// One [`ZeroPadMleCheckProver`] does that for both.
64pub(super) struct PaddedLayerProver<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
65	ZeroPadMleCheckProver<F, TreeLayer<'a, A, F, P>>,
66);
67
68impl<A, F, P> MleCheckProver<F> for PaddedLayerProver<'_, A, F, P>
69where
70	A: Allocator,
71	F: Field,
72	P: PackedField<Scalar = F>,
73{
74	fn n_vars(&self) -> usize {
75		self.0.n_vars()
76	}
77
78	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
79		self.0.execute()
80	}
81
82	fn fold(&mut self, challenge: F) {
83		self.0.fold(challenge);
84	}
85
86	fn finish(self) -> Vec<F> {
87		self.0.finish()
88	}
89
90	fn eval_point(&self) -> &[F] {
91		self.0.eval_point()
92	}
93}
94
95/// A batch's trees, lifted to one uniform depth by zero-fraction padding.
96///
97/// The batch runs a single layer schedule, so every tree in it must have the same depth.
98/// Padding lifts each tree to the depth of the deepest one.
99///
100/// A padded tree sits out the layers above its own root.
101/// It contributes a stand-in layer for each of them, and only then starts spending its real ones.
102///
103/// No padded layer is ever materialized.
104/// The whole padding is these two vectors, plus the $O(1)$ per-round correction each layer carries.
105pub(super) struct PaddedBatch<'a, A: Allocator, P: PackedField> {
106	/// The trees, in input order, each spending its layers deepest-first.
107	trees: Vec<FracAddCircuit<'a, A, P>>,
108	/// How much depth each tree is padded by, in tree order.
109	pad_lens: Vec<usize>,
110	/// The depth every tree is padded to, which is how many layers the schedule runs.
111	n_layers: usize,
112	/// Scratch space for the inversion that de-pads a layer's claims.
113	///
114	/// Every layer inverts one weight per tree, so the size never changes.
115	pad_eq_inversion: BatchInversion<P::Scalar>,
116}
117
118impl<'a, A, F, P> PaddedBatch<'a, A, P>
119where
120	A: Allocator,
121	F: Field,
122	P: PackedField<Scalar = F>,
123{
124	/// Pads `trees` up to the depth of the deepest one.
125	///
126	/// # Preconditions
127	/// * `trees` is non-empty
128	/// * at least one tree has a layer
129	pub(super) fn new(trees: Vec<FracAddCircuit<'a, A, P>>) -> Self {
130		let n_layers = trees
131			.iter()
132			.map(FracAddCircuit::n_layers)
133			.max()
134			.expect("precondition: trees is non-empty");
135		assert!(n_layers >= 1); // precondition
136
137		let pad_lens = trees
138			.iter()
139			.map(|tree| n_layers - tree.n_layers())
140			.collect();
141
142		let pad_eq_inversion = BatchInversion::new(trees.len());
143
144		Self {
145			trees,
146			pad_lens,
147			n_layers,
148			pad_eq_inversion,
149		}
150	}
151
152	/// How many layers the batch's schedule runs.
153	pub(super) const fn n_layers(&self) -> usize {
154		self.n_layers
155	}
156
157	/// How many trees the batch holds.
158	pub(super) const fn n_trees(&self) -> usize {
159		self.trees.len()
160	}
161
162	/// The allocator every tree's layer buffers are drawn from.
163	pub(super) fn alloc(&self) -> &'a A {
164		self.trees[0].alloc
165	}
166
167	/// Consumes the spent batch.
168	///
169	/// A depth-zero tree is all padding, so it is passed through every layer and never popped.
170	/// Every tree that had layers has spent them all.
171	pub(super) fn finish(self) {
172		debug_assert!(
173			self.trees.iter().all(|tree| tree.n_layers() == 0),
174			"every tree with layers is exhausted after n_layers reductions"
175		);
176	}
177
178	/// Pops one padded layer prover per tree, for the layer claimed at `node_point`.
179	///
180	/// The provers come back in tree order, so they stay aligned with `claims`.
181	///
182	/// A tree the batch has not reached yet keeps every layer it has.
183	pub(super) fn pop_layer(
184		&mut self,
185		claims: &[Fraction<F>],
186		node_point: &[F],
187	) -> Vec<PaddedLayerProver<'a, A, F, P>> {
188		let node_len = node_point.len();
189
190		// Every tree's padding segment is a prefix of this one node point, so a single table of
191		// prefix products serves the whole batch.
192		let pad_eq_prefixes = iter::once(F::ONE)
193			.chain(node_point.iter().scan(F::ONE, |acc, &coord| {
194				*acc *= eq_one_var(F::ZERO, coord);
195				Some(*acc)
196			}))
197			.collect::<Vec<_>>();
198
199		// De-padding a claim divides by the padding segment's equality weight, so the batch pays
200		// one inversion rather than one per tree.
201		let mut pad_eq_invs = self
202			.pad_lens
203			.iter()
204			.map(|&pad_len| pad_eq_prefixes[pad_len.min(node_len)])
205			.collect::<Vec<_>>();
206		assert!(
207			pad_eq_invs.iter().all(|&pad_eq| pad_eq != F::ZERO),
208			"a padding coordinate of the claim point equals one"
209		);
210		self.pad_eq_inversion.invert_nonzero(&mut pad_eq_invs);
211
212		izip!(&mut self.trees, &self.pad_lens, claims, &pad_eq_invs)
213			.map(|(tree, &tree_pad_len, &Fraction { num, den }, &pad_eq_inv)| {
214				let pad_len = tree_pad_len.min(node_len);
215				let point = node_point[pad_len..].to_vec();
216				let [num_claim, den_claim] = zero_pad_mle::unpad_claims(pad_eq_inv, [num, den]);
217
218				let inner = if node_len < tree_pad_len {
219					// The batch is still above this tree, so every variable of its layer is a
220					// padding variable and the de-padded claim is the tree's own fractional sum.
221					// The layer is that fraction beside the zero fraction 0/1, and the tree keeps
222					// all of its layers.
223					Either::Right(ConstantFraction::new(num_claim, den_claim))
224				} else {
225					Either::Left(tree.pop_layer(FracAddEvalClaim {
226						num_eval: num_claim,
227						den_eval: den_claim,
228						point,
229					}))
230				};
231
232				PaddedLayerProver(zero_pad_mle::new(
233					pad_eq_prefixes[..=pad_len].to_vec(),
234					node_point.to_vec(),
235					inner,
236				))
237			})
238			.collect()
239	}
240}
241
242/// Reduces a leaf claim on a zero-fraction-padded witness to the claim on the witness itself.
243///
244/// A batched fractional-addition check over trees of unequal depths pads each shallow tree.
245/// [`binius_ip::fracaddcheck::verify`] is oblivious to that padding.
246/// So the claims it outputs for such a tree are claims on the padded witness $(N', D')$.
247/// Its padding variables are the lowest `n_pad_vars` coordinates of `point`.
248///
249/// This divides out the padding variables' equality weight and drops them from the point.
250/// What remains are the claims on $N$ and $D$.
251///
252/// # Arguments
253///
254/// * `fraction` - The claimed numerator and denominator evaluations of the padded witness.
255/// * `point` - The reduced evaluation point, with the batch's selector coordinates already
256///   stripped.
257/// * `n_pad_vars` - How much depth this tree was padded by: the batch's layer count less the tree's
258///   own.
259///
260/// # Preconditions
261/// * `point.len() >= n_pad_vars`
262///
263/// # Panics
264///
265/// Panics if the padding coordinates' equality weight is zero, which requires one of them to equal
266/// one. They are the verifier's own challenges, so no prover can induce this; it happens with
267/// probability at most $\nu / |K|$.
268pub fn unpad_leaf_claim<F: Field>(
269	fraction: Fraction<F>,
270	point: &[F],
271	n_pad_vars: usize,
272) -> FracAddEvalClaim<F> {
273	assert!(point.len() >= n_pad_vars); // precondition
274
275	let pad_eq = point[..n_pad_vars]
276		.iter()
277		.map(|&coord| eq_one_var(F::ZERO, coord))
278		.product::<F>();
279	assert!(pad_eq != F::ZERO, "a padding coordinate equals one");
280	let pad_eq_inv = pad_eq.invert_or_zero();
281
282	let Fraction {
283		num: num_eval,
284		den: den_eval,
285	} = fraction;
286	FracAddEvalClaim {
287		num_eval: num_eval * pad_eq_inv,
288		den_eval: F::ONE + (den_eval - F::ONE) * pad_eq_inv,
289		point: point[n_pad_vars..].to_vec(),
290	}
291}
292
293#[cfg(test)]
294mod tests {
295	use binius_field::FieldOps;
296	use binius_ip::fracaddcheck;
297	use binius_math::test_utils::{Packed128b, random_scalars};
298	use proptest::prelude::*;
299	use rand::prelude::*;
300
301	use super::*;
302
303	type F = <Packed128b as FieldOps>::Scalar;
304
305	proptest! {
306		// Invariant: `unpad_leaf_claim` is the exact inverse of `pad_leaf_fraction`.
307		//
308		// The verifier pads a transparent leaf fraction, the prover unpads the claim it gets back.
309		// Only an end-to-end proof failure notices if either map drifts from the other.
310		#[test]
311		fn unpad_leaf_claim_inverts_pad_leaf_fraction(
312			seed in any::<u64>(),
313			n_pad_vars in 0usize..=5,
314			n_real_vars in 0usize..=5,
315		) {
316			let mut rng = StdRng::seed_from_u64(seed);
317
318			// Splitting the point's length in two keeps `n_pad_vars <= point.len()` by construction.
319			let point = random_scalars::<F>(&mut rng, n_pad_vars + n_real_vars);
320			let halves = random_scalars::<F>(&mut rng, 2);
321			let fraction = Fraction::new(halves[0], halves[1]);
322
323			let pad_eq = point[..n_pad_vars]
324				.iter()
325				.map(|&coord| eq_one_var(F::ZERO, coord))
326				.product::<F>();
327			// Unpadding asserts on a zero weight, which needs a padding coordinate equal to one.
328			// Random 128-bit coordinates never are, so this rejects nothing.
329			prop_assume!(pad_eq != F::ZERO);
330
331			let padded = fracaddcheck::pad_leaf_fraction(fraction.into(), pad_eq);
332			let claim = unpad_leaf_claim(padded.into(), &point, n_pad_vars);
333
334			prop_assert_eq!(claim.num_eval, fraction.num);
335			prop_assert_eq!(claim.den_eval, fraction.den);
336			// The padding variables are the lowest ones, so unpadding strips them off the point.
337			prop_assert_eq!(claim.point, point[n_pad_vars..].to_vec());
338		}
339	}
340}