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}