Skip to main content

binius_ip_prover/fracaddcheck/
circuit.rs

1// Copyright 2025-2026 The Binius Developers
2
3//! The materialized layers of a fractional-addition circuit, and the driver that proves them.
4
5use binius_compute::{Allocator, VecLike};
6use binius_field::{Field, PackedField};
7use binius_ip::fracaddcheck::FracAddEvalClaim;
8use binius_math::{FieldBuffer, FieldVec, line::extrapolate_line};
9use binius_utils::rayon::{
10	iter::{IntoParallelIterator, ParallelIterator},
11	task_size::{IndexedParallelIteratorExt, WorkPerItem},
12};
13
14use super::{LayerProver, fraction::Fraction};
15use crate::{
16	channel::IPProverChannel,
17	sumcheck::{batch::batch_prove_mle, frac_add_mle},
18};
19
20/// The materialized layers of a fractional addition circuit.
21///
22/// Each layer holds the numerator and denominator of one fractional term per node.
23/// A layer is half the width of the one below it, each node adding its two children:
24/// $$\frac{a_0}{b_0} + \frac{a_1}{b_1} = \frac{a_0b_1 + a_1b_0}{b_0b_1}$$
25pub struct FracAddCircuit<'a, A: Allocator, P: PackedField> {
26	layers: Vec<Fraction<FieldVec<P, A>>>,
27	/// Allocator the layer buffers are drawn from.
28	pub(crate) alloc: &'a A,
29}
30
31impl<A: Allocator, P: PackedField> Clone for FracAddCircuit<'_, A, P>
32where
33	A::Vec<P>: Clone,
34{
35	fn clone(&self) -> Self {
36		Self {
37			layers: self.layers.clone(),
38			alloc: self.alloc,
39		}
40	}
41}
42
43impl<'a, A, F, P> FracAddCircuit<'a, A, P>
44where
45	A: Allocator,
46	F: Field,
47	P: PackedField<Scalar = F>,
48{
49	/// Materializes every layer of the circuit from `witness`.
50	///
51	/// Returns the circuit beside `sums`, its root layer.
52	/// `sums` holds the fractional addition over all `k` reduced variables.
53	///
54	/// # Arguments
55	/// * `k` - How many variables to reduce, one sibling fractional addition per step.
56	/// * `witness` - The witness numerator/denominator layers
57	///
58	/// # Preconditions
59	/// * `witness.num.log_len() >= k`
60	pub fn build(
61		k: usize,
62		alloc: &'a A,
63		witness: Fraction<FieldVec<P, A>>,
64	) -> (Self, Fraction<FieldVec<P, A>>) {
65		let Fraction {
66			num: witness_num,
67			den: witness_den,
68		} = witness;
69		assert_eq!(
70			witness_num.log_len(),
71			witness_den.log_len(),
72			"numerator and denominator witnesses must have equal length"
73		);
74		assert!(witness_num.log_len() >= k);
75
76		let mut layers = Vec::with_capacity(k + 1);
77		layers.push(Fraction::new(witness_num, witness_den));
78
79		for _ in 0..k {
80			let prev_layer = layers.last().expect("layers is non-empty");
81
82			let Fraction { num, den } = prev_layer;
83			let num_log_len = num.log_len() - 1;
84			let den_log_len = den.log_len() - 1;
85			let (num_0, num_1) = num.split_half();
86			let (den_0, den_1) = den.split_half();
87
88			// One packed word of the next layer from the sibling halves, written straight into
89			// the pooled buffers:
90			//     a_0/b_0 + a_1/b_1 = (a_0*b_1 + a_1*b_0) / (b_0*b_1)
91			// Workers each take a contiguous run of words.
92			// One word is three multiplies and an add, a few nanoseconds of work.
93			// A run must therefore be long enough to pay back handing it off.
94			let out_len = num_0.as_ref().len();
95			let mut num_data = alloc.alloc::<P>(out_len);
96			let mut den_data = alloc.alloc::<P>(out_len);
97			(
98				num_data.spare_capacity_mut(),
99				den_data.spare_capacity_mut(),
100				num_0.as_ref(),
101				den_0.as_ref(),
102				num_1.as_ref(),
103				den_1.as_ref(),
104			)
105				.into_par_iter()
106				.with_min_task(WorkPerItem::FieldMuls)
107				.for_each(|(num_out, den_out, &num_0, &den_0, &num_1, &den_1)| {
108					num_out.write(num_0 * den_1 + num_1 * den_0);
109					den_out.write(den_0 * den_1);
110				});
111			// Invariant: every zip input holds at least `out_len` words.
112			//
113			// A parallel zip yields as many items as its shortest input holds.
114			// A shorter input would leave trailing slots uninitialized.
115			//
116			//     spare capacity:  >= out_len   allocated for at least that many
117			//     sibling halves:  == out_len   halves of two equal-length buffers
118			assert!(
119				num_data.capacity() - num_data.len() >= out_len
120					&& den_data.capacity() - den_data.len() >= out_len,
121				"allocated buffers must hold every claimed slot"
122			);
123			assert!(
124				[den_0.as_ref(), num_1.as_ref(), den_1.as_ref()]
125					.iter()
126					.all(|half| half.len() == out_len),
127				"the four sibling halves must hold exactly one word per claimed slot"
128			);
129			// Safety: both length claims cover only initialized slots.
130			// - The assertions above bound every zip input below by `out_len`.
131			// - So the loop ran `out_len` items.
132			// - Each item wrote one numerator slot and one denominator slot.
133			unsafe {
134				num_data.set_len(out_len);
135				den_data.set_len(out_len);
136			}
137			let next_layer = Fraction::new(
138				FieldBuffer::new(num_log_len, num_data),
139				FieldBuffer::new(den_log_len, den_data),
140			);
141
142			layers.push(next_layer);
143		}
144
145		let sums = layers.pop().expect("layers has k+1 elements");
146		(Self { layers, alloc }, sums)
147	}
148
149	/// Returns the number of remaining layers to prove.
150	pub const fn n_layers(&self) -> usize {
151		self.layers.len()
152	}
153
154	/// Pops the widest remaining layer as the MLE-check prover that reduces it.
155	///
156	/// The returned prover owns the popped buffers and borrows only the allocator.
157	/// So it outlives this borrow, and the circuit stays in place while a caller drives it.
158	///
159	/// # Preconditions
160	/// * `self.n_layers() >= 1`
161	pub fn pop_layer(&mut self, claim: FracAddEvalClaim<F>) -> LayerProver<'a, A, F, P> {
162		let Fraction { num, den } = self
163			.layers
164			.pop()
165			.expect("precondition: self.n_layers() >= 1");
166
167		// The MLE-check reduces four multilinears: the low and high halves of the numerator buffer
168		// and of the denominator buffer. The store takes ownership of the two popped buffers and
169		// shares each between its halves, so the prover is self-contained with no up-front copy of
170		// the popped layer.
171		frac_add_mle::new_split_half(
172			self.alloc,
173			num,
174			den,
175			claim.point,
176			[claim.num_eval, claim.den_eval],
177		)
178	}
179
180	/// Runs the fractional addition check protocol and returns the final evaluation claims.
181	///
182	/// This consumes the circuit, reducing from the smallest layer back to the largest.
183	///
184	/// # Arguments
185	/// * `claim` - The numerator and denominator claims at their shared evaluation point.
186	/// * `channel` - The channel for sending prover messages and sampling challenges.
187	///
188	/// # Preconditions
189	/// * `claim.point.len() == witness.log_len() - k`, for `k` the number of reduction layers.
190	pub fn prove(
191		self,
192		claim: FracAddEvalClaim<F>,
193		channel: &mut impl IPProverChannel<F>,
194	) -> FracAddEvalClaim<F> {
195		// Proving the full circuit runs every layer, so delegate and drop the leftover circuit.
196		let n_layers = self.n_layers();
197		let (remaining, claim) = self.prove_layers(n_layers, claim, channel);
198		debug_assert_eq!(remaining.n_layers(), 0, "proving every layer leaves none unproved");
199		claim
200	}
201
202	/// Runs the first `n_layers` fractional-addition layers from a claim, returning the remainder.
203	///
204	/// Each layer adds one variable via a sumcheck and a line-fold.
205	/// So starting from a claim over `d` variables, the returned claim is over `d + n_layers`.
206	///
207	/// This is the layer loop of [`Self::prove`], which runs every layer.
208	/// The returned circuit still holds the layers that were not proved.
209	///
210	/// # Arguments
211	/// * `n_layers` - The number of layers to prove, at most [`Self::n_layers`].
212	/// * `claim` - The numerator and denominator claims at their shared evaluation point.
213	/// * `channel` - The channel for sending prover messages and sampling challenges.
214	///
215	/// # Returns
216	/// * the circuit, holding whatever layers were not proved,
217	/// * the reduced numerator/denominator claims after `n_layers` layers.
218	///
219	/// # Preconditions
220	/// * `n_layers <= self.n_layers()`.
221	fn prove_layers(
222		mut self,
223		n_layers: usize,
224		claim: FracAddEvalClaim<F>,
225		channel: &mut impl IPProverChannel<F>,
226	) -> (Self, FracAddEvalClaim<F>) {
227		let mut claim = claim;
228
229		for _ in 0..n_layers {
230			let sumcheck_prover = self.pop_layer(claim);
231
232			// The driver draws the batching coefficient and Horner-folds the layer's two claims,
233			// which is the polynomial the verifier's `batch_verify_mle` reconstructs.
234			let output = batch_prove_mle(vec![sumcheck_prover], channel);
235			output.send_evals(channel);
236
237			let mut multilinear_evals = output.multilinear_evals;
238			let evals = multilinear_evals.pop().expect("batch contains one prover");
239
240			let [num_0, num_1, den_0, den_1] = evals
241				.try_into()
242				.expect("prover evaluates four multilinears");
243
244			// Fold the highest variable to combine the two halves into the next layer's claim.
245			let r = channel.sample();
246
247			let next_num = extrapolate_line(num_0, num_1, r);
248			let next_den = extrapolate_line(den_0, den_1, r);
249
250			// Sumcheck binds variables high-to-low; reverse to low-to-high for the claim point.
251			let mut next_point = output.challenges;
252			next_point.reverse();
253			next_point.push(r);
254
255			claim = FracAddEvalClaim {
256				num_eval: next_num,
257				den_eval: next_den,
258				point: next_point,
259			};
260		}
261
262		(self, claim)
263	}
264}
265
266#[cfg(test)]
267mod tests {
268	use std::iter;
269
270	use binius_compute::GlobalAllocator;
271	use binius_ip::fracaddcheck;
272	use binius_math::{
273		multilinear::evaluate::evaluate,
274		test_utils::{Packed128b, random_field_buffer, random_scalars},
275	};
276	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
277	use proptest::prelude::*;
278	use rand::prelude::*;
279
280	use super::*;
281
282	type StdChallenger = HasherChallenger<sha2::Sha256>;
283
284	fn test_frac_add_check_prove_verify_helper<P: PackedField>(n: usize, k: usize) {
285		let mut rng = StdRng::seed_from_u64(0);
286		let alloc = GlobalAllocator;
287
288		// 1. Create random witness with log_len = n + k
289		let witness_num = random_field_buffer::<P>(&mut rng, n + k);
290		let witness_den = random_field_buffer::<P>(&mut rng, n + k);
291
292		// 2. Create prover (computes fractional-add layers)
293		let (prover, sums) = FracAddCircuit::build(
294			k,
295			&alloc,
296			Fraction::new(witness_num.clone(), witness_den.clone()),
297		);
298
299		// 3. Generate random n-dimensional challenge point
300		let eval_point = random_scalars::<P::Scalar>(&mut rng, n);
301
302		// 4. Evaluate sums at challenge point to create claims
303		let sum_num_eval = evaluate(&sums.num, &eval_point);
304		let sum_den_eval = evaluate(&sums.den, &eval_point);
305		// The prover and the verifier take the same claim type, so one claim serves both.
306		let claim = FracAddEvalClaim {
307			num_eval: sum_num_eval,
308			den_eval: sum_den_eval,
309			point: eval_point,
310		};
311
312		// 5. Run prover
313		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
314		let prover_output = prover.prove(claim.clone(), &mut prover_transcript);
315
316		// 6. Run verifier
317		let mut verifier_transcript = prover_transcript.into_verifier();
318		let verifier_output = fracaddcheck::verify(k, claim, &mut verifier_transcript).unwrap();
319
320		// 7. Check outputs match
321		assert_eq!(prover_output, verifier_output);
322
323		// 8. Verify multilinear evaluation of original witness
324		let expected_num = evaluate(&witness_num, &verifier_output.point);
325		let expected_den = evaluate(&witness_den, &verifier_output.point);
326		assert_eq!(verifier_output.num_eval, expected_num);
327		assert_eq!(verifier_output.den_eval, expected_den);
328	}
329
330	#[test]
331	fn test_frac_add_check_prove_verify() {
332		test_frac_add_check_prove_verify_helper::<Packed128b>(4, 3);
333	}
334
335	#[test]
336	fn test_frac_add_check_full_prove_verify() {
337		test_frac_add_check_prove_verify_helper::<Packed128b>(0, 4);
338	}
339
340	fn check_all_layers<P: PackedField>(n: usize, k: usize, seed: u64) {
341		let mut rng = StdRng::seed_from_u64(seed);
342		let alloc = GlobalAllocator;
343
344		// Create random witness with log_len = n + k
345		let witness_num = random_field_buffer::<P>(&mut rng, n + k);
346		let witness_den = random_field_buffer::<P>(&mut rng, n + k);
347
348		// Create prover (computes fractional-add layers)
349		let (prover, sums) = FracAddCircuit::build(
350			k,
351			&alloc,
352			Fraction::new(witness_num.clone(), witness_den.clone()),
353		);
354
355		// `build` pops the root off as `sums`, so the circuit is `layers` followed by it.
356		for (j, layer) in prover.layers.iter().chain(iter::once(&sums)).enumerate() {
357			// Entry i of layer j is the fractional sum of the 2^j witness values strided by that
358			// layer's own width (strided access, not contiguous).
359			let width = 1 << (n + k - j);
360			let num_terms = 1 << j;
361			for i in 0..width {
362				let mut expected_num = witness_num.get(i);
363				let mut expected_den = witness_den.get(i);
364				for z in 1..num_terms {
365					let idx = i + z * width;
366					let num_z = witness_num.get(idx);
367					let den_z = witness_den.get(idx);
368					expected_num = expected_num * den_z + num_z * expected_den;
369					expected_den *= den_z;
370				}
371				let actual_num = layer.num.get(i);
372				let actual_den = layer.den.get(i);
373				assert_eq!(actual_num, expected_num, "layer {j} numerator mismatch at index {i}");
374				assert_eq!(actual_den, expected_den, "layer {j} denominator mismatch at index {i}");
375			}
376		}
377	}
378
379	proptest! {
380		// Invariant: every layer of the circuit is the fractional-addition fold of the witness.
381		//
382		// Pinning each layer to that fold pins the sibling recurrence the layers are built from.
383		// Only an end-to-end proof failure notices if `build` folds the wrong pairs.
384		#[test]
385		fn frac_add_check_layers_fold_the_witness(
386			seed in any::<u64>(),
387			n in 0usize..=4,
388			k in 0usize..=4,
389		) {
390			check_all_layers::<Packed128b>(n, k, seed);
391		}
392	}
393}