Skip to main content

binius_ip_prover/fracaddcheck/
driver.rs

1// Copyright 2025-2026 The Binius Developers
2
3//! The batched layer schedule: one uniform round dance over every tree in the batch.
4
5use std::iter;
6
7use binius_compute::Allocator;
8use binius_field::{Field, PackedField};
9use binius_ip::{mlecheck, sumcheck::RoundCoeffs};
10use binius_math::{
11	FieldBuffer, FieldVec, line::extrapolate_line, multilinear::eq::eq_ind_partial_eval,
12};
13use binius_utils::rayon::iter::{IntoParallelRefMutIterator, ParallelIterator};
14use itertools::izip;
15
16use super::{FracAddCircuit, fraction::Fraction, padding::PaddedBatch};
17use crate::{
18	channel::IPProverChannel,
19	sumcheck::{
20		common::MleCheckProver,
21		frac_add_mle,
22		mle_store::MleStore,
23		round_evaluator::{MleCheckRoundEvaluator, SharedMleCheckProver},
24	},
25};
26
27/// Output of [`batch_prove_unequal_depths`].
28///
29/// After the full `n_layers` reduction, `fractions` holds each input tree's reduced fraction at
30/// `eval_point`. The batched claim the verifier checks is the eq(selector)-weighted combination of
31/// these fractions.
32pub struct BatchProveOutput<F> {
33	/// The reduced evaluation point (`selector ++ content`) at which the fractions are claimed.
34	pub eval_point: Vec<F>,
35	/// Each input prover's reduced `(num, den)` fraction at `eval_point`, in input order.
36	pub fractions: Vec<Fraction<F>>,
37}
38
39/// Runs one batched fracaddcheck layer given its per-instance final-layer MLE-check provers.
40///
41/// The layer runs in four steps, one function each:
42/// - [`prove_content_rounds`] folds the content variables of every instance in lockstep.
43/// - [`finish_and_transpose`] turns the reduced halves into the four selector columns.
44/// - [`prove_selector_rounds`] folds the `k` selector variables in one MLE-check.
45/// - [`finalize_layer`] line-folds the merged evaluations into the next layer's claims.
46///
47/// Returns the per-instance fractions and the next evaluation point.
48/// The fractions are padded to the `2^k` selector slots with the zero fraction.
49///
50/// One `batch_coeff` batches the layer's numerator and denominator claims.
51/// The verifier's `batch_verify_mle` samples it once per layer, before the round polynomials.
52/// The content rounds and the selector rounds reuse the same coefficient.
53fn reduce_layer<A, F, P, MP>(
54	alloc: &A,
55	mut layer_provers: Vec<MP>,
56	eval_point: &[F],
57	k: usize,
58	channel: &mut impl IPProverChannel<F>,
59) -> (Vec<Fraction<F>>, Vec<F>)
60where
61	A: Allocator,
62	F: Field,
63	P: PackedField<Scalar = F>,
64	MP: MleCheckProver<F> + Send,
65{
66	// Split eval_point into outer (selector) and inner (content) coordinates.
67	let (outer_coords, inner_coords) = eval_point.split_at(k);
68
69	// eq weights for batching over instances: eq(i, outer_coords) for all i in B_k.
70	let eq_weights = eq_ind_partial_eval::<F>(outer_coords);
71
72	let batch_coeff = channel.sample();
73
74	let mut challenges = Vec::with_capacity(eval_point.len());
75
76	prove_content_rounds(
77		&mut layer_provers,
78		eq_weights.as_ref(),
79		inner_coords.len(),
80		batch_coeff,
81		&mut challenges,
82		channel,
83	);
84
85	let (reduced_halves, selector_columns) =
86		finish_and_transpose::<A, F, P, MP>(alloc, layer_provers, k);
87
88	let merged_evals = prove_selector_rounds(
89		alloc,
90		selector_columns,
91		eq_weights.as_ref(),
92		outer_coords,
93		batch_coeff,
94		&mut challenges,
95		channel,
96	);
97
98	finalize_layer(merged_evals, &reduced_halves, k, challenges, channel)
99}
100
101/// Folds the content variables of every instance in lockstep, one round polynomial per round.
102///
103/// Each round sends the eq(selector)-weighted sum of the per-instance round polynomials.
104/// One instance's polynomial is its `[num, den]` pair batched with `batch_coeff`.
105/// Every instance then folds on the challenge the round draws.
106/// The challenges are appended to `challenges` in round order.
107///
108/// `eq_weights` holds one weight per selector slot, the instances taking the leading ones.
109/// A slot past the last instance holds the constant fraction 0/1: numerator 0, denominator 1.
110/// A constant composition has that same constant as its round polynomial.
111/// So a padding slot's claims stay (0, 1) through every fold.
112/// It contributes `eq_i * batch_coeff` to each round polynomial's constant coefficient.
113fn prove_content_rounds<F, MP>(
114	layer_provers: &mut [MP],
115	eq_weights: &[F],
116	n_rounds: usize,
117	batch_coeff: F,
118	challenges: &mut Vec<F>,
119	channel: &mut impl IPProverChannel<F>,
120) where
121	F: Field,
122	MP: MleCheckProver<F> + Send,
123{
124	let pad_eq_sum: F = eq_weights[layer_provers.len()..].iter().copied().sum();
125
126	for _round in 0..n_rounds {
127		// The instances are independent within a round, so their polynomials compute in parallel.
128		//
129		// One instance's round is too small a parallel region to fill the pool alone.
130		let per_instance: Vec<RoundCoeffs<F>> = layer_provers
131			.par_iter_mut()
132			.map(|prover| RoundCoeffs::batch(prover.execute(), &batch_coeff))
133			.collect();
134
135		// Weight instance j's polynomial by eq_j and sum, in instance order.
136		let real_coeffs: RoundCoeffs<F> = iter::zip(per_instance, eq_weights)
137			.map(|(coeffs, &eq_i)| coeffs * eq_i)
138			.sum();
139		let round_coeffs = real_coeffs + &RoundCoeffs(vec![pad_eq_sum * batch_coeff]);
140
141		channel.send_many(mlecheck::RoundProof::truncate(round_coeffs).coeffs());
142
143		let challenge = channel.sample();
144		challenges.push(challenge);
145
146		for prover in layer_provers.iter_mut() {
147			prover.fold(challenge);
148		}
149	}
150}
151
152/// Finishes the content provers and transposes their reduced halves into the selector columns.
153///
154/// Each instance finishes with the four evaluations `[num_0, num_1, den_0, den_1]` it reduced.
155/// Instance `i` occupies slot `i` of each of the four returned columns.
156/// The columns span the `k` selector variables, so they hold `2^k` slots.
157///
158/// Returns the per-instance evaluations and those columns.
159/// The line-fold that closes the layer reduces the evaluations.
160/// The selector MLE-check folds the columns.
161///
162/// Both children of a padding slot are the zero fraction 0/1.
163/// So a slot past the last instance holds 0 in the numerator columns and 1 in the denominators.
164fn finish_and_transpose<A, F, P, MP>(
165	alloc: &A,
166	layer_provers: Vec<MP>,
167	k: usize,
168) -> (Vec<[F; 4]>, [FieldVec<P, A>; 4])
169where
170	A: Allocator,
171	F: Field,
172	P: PackedField<Scalar = F>,
173	MP: MleCheckProver<F>,
174{
175	let reduced: Vec<[F; 4]> = layer_provers
176		.into_iter()
177		.map(|prover| {
178			prover
179				.finish()
180				.try_into()
181				.expect("fractional-addition prover has four multilinears")
182		})
183		.collect();
184
185	// Each column starts as all padding, then one pass over the reduced halves writes the
186	// instances into its leading slots.
187	let pad = Fraction::<F>::ZERO;
188	let mut columns = [pad.num, pad.num, pad.den, pad.den].map(|pad_half| {
189		let mut column = FieldBuffer::zeros_in(alloc, k);
190		for slot in reduced.len()..1 << k {
191			column.set(slot, pad_half);
192		}
193		column
194	});
195	for (slot, evals) in reduced.iter().enumerate() {
196		for (column, &eval) in iter::zip(&mut columns, evals) {
197			column.set(slot, eval);
198		}
199	}
200
201	(reduced, columns)
202}
203
204/// Folds the `k` selector variables of one layer in a single fractional-addition MLE-check.
205///
206/// The claim is what the content rounds reduced the layer to.
207/// It is the eq(selector)-weighted sum of the fractional-addition composition of the four columns.
208/// The rounds reuse `batch_coeff`, and their challenges are appended to `challenges`.
209///
210/// Returns the merged `[num_0, num_1, den_0, den_1]` evaluations at those challenges.
211fn prove_selector_rounds<'a, A, F, P>(
212	alloc: &'a A,
213	columns: [FieldVec<P, A>; 4],
214	eq_weights: &[F],
215	outer_coords: &[F],
216	batch_coeff: F,
217	challenges: &mut Vec<F>,
218	channel: &mut impl IPProverChannel<F>,
219) -> [F; 4]
220where
221	A: Allocator,
222	F: Field,
223	P: PackedField<Scalar = F>,
224{
225	let k = outer_coords.len();
226
227	let [num_0s, num_1s, den_0s, den_1s] = &columns;
228	let num_eval: F = izip!(
229		num_0s.iter_scalars(),
230		num_1s.iter_scalars(),
231		den_0s.iter_scalars(),
232		den_1s.iter_scalars(),
233		eq_weights
234	)
235	.map(|(n0, n1, d0, d1, &eq_i)| eq_i * (n0 * d1 + n1 * d0))
236	.sum();
237	let den_eval: F = izip!(den_0s.iter_scalars(), den_1s.iter_scalars(), eq_weights)
238		.map(|(d0, d1, &eq_i)| eq_i * (d0 * d1))
239		.sum();
240
241	// The columns are freshly packed for this check, so the store owns them directly.
242	let mut selector_store = MleStore::new(k, alloc);
243	let selector_cols = columns.map(|column| selector_store.push_owned(column));
244	let (selector_num, selector_den) = frac_add_mle::evaluators::<F, P>(selector_cols);
245	let claims_with_evaluators: [(F, Box<dyn MleCheckRoundEvaluator<F, P> + 'a>); 2] = [
246		(num_eval, Box::new(selector_num)),
247		(den_eval, Box::new(selector_den)),
248	];
249	let mut selector_prover =
250		SharedMleCheckProver::new(selector_store, claims_with_evaluators, outer_coords.to_vec());
251
252	for _round in 0..k {
253		let round_coeffs = RoundCoeffs::batch(selector_prover.execute(), &batch_coeff);
254		channel.send_many(mlecheck::RoundProof::truncate(round_coeffs).coeffs());
255
256		let challenge = channel.sample();
257		challenges.push(challenge);
258		selector_prover.fold(challenge);
259	}
260
261	selector_prover
262		.finish()
263		.try_into()
264		.expect("fractional-addition prover has four multilinears")
265}
266
267/// Sends the merged child evaluations and line-folds them into the next layer's claims.
268///
269/// The verifier reads the four evaluations before it samples the doubling coordinate `r`.
270/// So the fold happens only once the transcript holds them.
271///
272/// `reduced` holds the four child evaluations of each real instance, one entry per instance.
273/// The `2^k - reduced.len()` padding slots hold the zero fraction.
274/// A line-fold between two zero fractions leaves it unchanged.
275/// Returning them keeps the output aligned with the next layer's `2^k` selector eq weights.
276fn finalize_layer<F: Field>(
277	merged_evals: [F; 4],
278	reduced: &[[F; 4]],
279	k: usize,
280	challenges: Vec<F>,
281	channel: &mut impl IPProverChannel<F>,
282) -> (Vec<Fraction<F>>, Vec<F>) {
283	channel.send_many(&merged_evals);
284
285	let r = channel.sample();
286
287	// Sumcheck binds variables high-to-low; reverse to low-to-high for the claim point.
288	let mut next_point = challenges;
289	next_point.reverse();
290	next_point.push(r);
291
292	let next_fractions = reduced
293		.iter()
294		.map(|&[num_0, num_1, den_0, den_1]| {
295			Fraction::new(extrapolate_line(num_0, num_1, r), extrapolate_line(den_0, den_1, r))
296		})
297		.chain(iter::repeat_n(Fraction::ZERO, (1 << k) - reduced.len()))
298		.collect();
299
300	(next_fractions, next_point)
301}
302
303/// Runs a batched fractional-addition check for trees of *unequal* depths.
304///
305/// Every tree shallower than the deepest is proved over its zero-fraction-padded witness.
306/// [`super::padding`] states the identity that such a padding satisfies.
307/// The transcript is then exactly that of an equal-depth batch of the maximum depth.
308/// The verifier runs the ordinary [`binius_ip::fracaddcheck::verify`] over `n_layers` layers.
309/// It never learns the individual depths.
310///
311/// Every prover must reduce over *all* of its witness variables, so each fractional sum is a
312/// scalar and there is no content point.
313/// Dropping the content dimension keeps the padding bookkeeping to four scalars per layer.
314///
315/// The prover does not materialize the padded witnesses.
316/// Each layer's per-tree reduction corrects the unpadded layer's messages in $O(1)$ per round.
317///
318/// # Arguments
319///
320/// * `provers` - The trees to batch, whose layer counts may differ.
321/// * `claimed_fractions` - Each tree's claimed root fraction, one per prover.
322/// * `selector_point` - Evaluation point for the selector variables.
323/// * `channel` - The channel for sending prover messages and sampling challenges.
324///
325/// # Preconditions
326/// * `provers` must be non-empty.
327/// * Every prover's witness must have exactly `prover.n_layers()` variables. A tree of depth zero
328///   is allowed — it is all padding, so its leaf claim is its root — but at least one tree must
329///   have a layer.
330/// * `2^selector_point.len() >= provers.len()`.
331/// * `claimed_fractions.len() == provers.len()`.
332///
333/// # Returns
334///
335/// A [`BatchProveOutput`] whose `fractions` are each tree's leaf claim, in input order, at the
336/// shared reduced `eval_point`.
337///
338/// Those leaf claims are on the *padded* witnesses.
339/// [`super::unpad_leaf_claim`] reduces one to the claims on the tree's own witness.
340pub fn batch_prove_unequal_depths<'a, A, F, P>(
341	provers: Vec<FracAddCircuit<'a, A, P>>,
342	claimed_fractions: Vec<Fraction<F>>,
343	selector_point: Vec<F>,
344	channel: &mut impl IPProverChannel<F>,
345) -> BatchProveOutput<F>
346where
347	A: Allocator,
348	F: Field,
349	P: PackedField<Scalar = F>,
350{
351	assert!(!provers.is_empty()); // precondition
352	assert_eq!(claimed_fractions.len(), provers.len()); // precondition
353
354	let k = selector_point.len();
355	assert!(provers.len() <= (1 << k)); // precondition
356
357	let mut batch = PaddedBatch::new(provers);
358	let alloc = batch.alloc();
359	let n_trees = batch.n_trees();
360
361	let mut claims = claimed_fractions;
362	let mut eval_point = selector_point;
363
364	// Each iteration reduces the layer whose node variables are the point's suffix past the
365	// selector coordinates. A tree the batch has not yet reached contributes a padding layer.
366	for _ in 0..batch.n_layers() {
367		let layer_provers = batch.pop_layer(&claims, &eval_point[k..]);
368		let (next_claims, next_point) =
369			reduce_layer::<A, F, P, _>(alloc, layer_provers, &eval_point, k, channel);
370		claims = next_claims;
371		eval_point = next_point;
372	}
373	batch.finish();
374
375	// `reduce_layer` pads its output to the 2^k selector slots; only the real trees remain.
376	let mut fractions = claims;
377	fractions.truncate(n_trees);
378
379	BatchProveOutput {
380		eval_point,
381		fractions,
382	}
383}
384
385#[cfg(test)]
386mod tests {
387	use binius_compute::GlobalAllocator;
388	use binius_ip::fracaddcheck;
389	use binius_math::{
390		inner_product::inner_product,
391		multilinear::evaluate::evaluate,
392		test_utils::{Packed128b, random_field_buffer, random_scalars},
393	};
394	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
395	use binius_utils::checked_arithmetics::log2_ceil_usize;
396	use proptest::prelude::*;
397	use rand::prelude::*;
398
399	use super::*;
400	use crate::fracaddcheck::unpad_leaf_claim;
401
402	type StdChallenger = HasherChallenger<sha2::Sha256>;
403
404	/// A numerator/denominator witness pair.
405	type Witness<P> = Fraction<FieldBuffer<P>>;
406
407	/// One prover per entry of `depths`, each reducing over all of its witness variables.
408	#[allow(clippy::type_complexity)]
409	fn unequal_depth_provers<'a, P: PackedField>(
410		rng: &mut impl rand::Rng,
411		alloc: &'a GlobalAllocator,
412		depths: &[usize],
413	) -> (Vec<Witness<P>>, Vec<FracAddCircuit<'a, GlobalAllocator, P>>, Vec<Fraction<P::Scalar>>) {
414		itertools::multiunzip(depths.iter().map(|&depth| {
415			let witness = Fraction::new(
416				random_field_buffer::<P>(&mut *rng, depth),
417				random_field_buffer::<P>(&mut *rng, depth),
418			);
419			let (prover, sums) = FracAddCircuit::build(depth, alloc, witness.clone());
420			assert_eq!(sums.num.log_len(), 0);
421			(witness, prover, sums.as_ref().map(|buffer| buffer.get(0)))
422		}))
423	}
424
425	/// The eq(selector)-weighted combination of per-tree fractions, as the verifier forms it.
426	///
427	/// The selector slots beyond the trees hold the zero fraction 0/1.
428	fn combine_fractions<P: PackedField>(
429		fractions: &[Fraction<P::Scalar>],
430		selector_point: &[P::Scalar],
431	) -> (P::Scalar, P::Scalar) {
432		let n_slots = 1 << selector_point.len();
433		let eq_weights = eq_ind_partial_eval::<P>(selector_point);
434		let num_eval = inner_product(
435			fractions.iter().map(|f| f.num),
436			(0..fractions.len()).map(|i| eq_weights.get(i)),
437		);
438		let den_eval = inner_product(
439			fractions
440				.iter()
441				.map(|f| f.den)
442				.chain(iter::repeat_n(P::Scalar::ONE, n_slots - fractions.len())),
443			(0..n_slots).map(|i| eq_weights.get(i)),
444		);
445		(num_eval, den_eval)
446	}
447
448	/// Proves a batch of unequal-depth trees against the depth-oblivious verifier, then unpads each
449	/// tree's leaf claims and checks them against that tree's own witness.
450	fn test_unequal_depths_helper<P: PackedField>(depths: &[usize], seed: u64) {
451		let mut rng = StdRng::seed_from_u64(seed);
452		let alloc = GlobalAllocator;
453
454		let k = log2_ceil_usize(depths.len());
455		let n_layers = *depths.iter().max().expect("depths is non-empty");
456
457		let (witnesses, provers, claimed_fractions) =
458			unequal_depth_provers::<P>(&mut rng, &alloc, depths);
459
460		// The verifier's input claim is the eq(selector)-weighted combination of the fractions.
461		let selector_point = random_scalars::<P::Scalar>(&mut rng, k);
462		let (num_eval, den_eval) = combine_fractions::<P>(&claimed_fractions, &selector_point);
463		let claim = fracaddcheck::FracAddEvalClaim {
464			num_eval,
465			den_eval,
466			point: selector_point.clone(),
467		};
468
469		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
470		let BatchProveOutput {
471			eval_point,
472			fractions,
473		} = batch_prove_unequal_depths(
474			provers,
475			claimed_fractions,
476			selector_point,
477			&mut prover_transcript,
478		);
479
480		// The verifier's control flow depends only on the maximum depth.
481		let mut verifier_transcript = prover_transcript.into_verifier();
482		let verifier_output =
483			fracaddcheck::verify(n_layers, claim, &mut verifier_transcript).unwrap();
484
485		assert_eq!(verifier_output.point, eval_point);
486		let (num_eval, den_eval) = combine_fractions::<P>(&fractions, &eval_point[..k]);
487		assert_eq!(verifier_output.num_eval, num_eval);
488		assert_eq!(verifier_output.den_eval, den_eval);
489
490		// Each tree's reduced claims are on its *padded* witness; unpadding them yields claims on
491		// the witness itself, at a suffix of the shared node point.
492		for (i, (&depth, witness)) in iter::zip(depths, &witnesses).enumerate() {
493			let leaf = unpad_leaf_claim(fractions[i], &eval_point[k..], n_layers - depth);
494			assert_eq!(leaf.point.len(), depth);
495			assert_eq!(leaf.num_eval, evaluate(&witness.num, &leaf.point), "tree {i} numerator");
496			assert_eq!(leaf.den_eval, evaluate(&witness.den, &leaf.point), "tree {i} denominator");
497		}
498	}
499
500	#[test]
501	fn test_unequal_depths_mixed() {
502		test_unequal_depths_helper::<Packed128b>(&[2, 4, 5], 11);
503	}
504
505	#[test]
506	fn test_unequal_depths_single_prover() {
507		test_unequal_depths_helper::<Packed128b>(&[3], 11);
508	}
509
510	#[test]
511	fn test_unequal_depths_power_of_two_provers() {
512		// The shallowest tree is padded by more than one layer, the deepest not at all.
513		test_unequal_depths_helper::<Packed128b>(&[1, 2, 5, 5], 11);
514	}
515
516	#[test]
517	fn test_unequal_depths_all_minimal() {
518		// Depth 1 throughout: every tree retains its final layer immediately.
519		test_unequal_depths_helper::<Packed128b>(&[1, 1, 1], 11);
520	}
521
522	#[test]
523	fn test_unequal_depths_zero_depth_tree() {
524		// A depth-0 tree never pops a layer: it is all padding, so its leaf claim is its root.
525		test_unequal_depths_helper::<Packed128b>(&[0, 3], 11);
526	}
527
528	#[test]
529	fn test_unequal_depths_maximal_padding() {
530		// A single-layer tree beside a deep one: all but its last reduction is padding.
531		test_unequal_depths_helper::<Packed128b>(&[1, 6], 11);
532	}
533
534	#[test]
535	fn test_unequal_depths_equal_depths() {
536		// Equal depths pad nothing, so every wrapper is a pass-through.
537		test_unequal_depths_helper::<Packed128b>(&[4, 4, 4], 11);
538	}
539
540	proptest! {
541		// A full batched prove-verify per case, so trade the default case count down for runtime.
542		#![proptest_config(ProptestConfig::with_cases(64))]
543
544		// Invariant: the batched round trip holds for any mix of tree depths.
545		//
546		// The cases above pin named edge shapes; this covers the space between them.
547		// Padding is bookkept per tree, so it is the mix of depths that stresses it.
548		#[test]
549		fn unequal_depths_round_trip(
550			seed in any::<u64>(),
551			depths in prop::collection::vec(0usize..=6, 1..=5),
552		) {
553			// Batching needs at least one layer to reduce, so an all-depth-0 batch is not a case.
554			prop_assume!(depths.iter().any(|&depth| depth > 0));
555
556			test_unequal_depths_helper::<Packed128b>(&depths, seed);
557		}
558	}
559}