binius_ip_prover/fracaddcheck/
circuit.rs1use 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
20pub struct FracAddCircuit<'a, A: Allocator, P: PackedField> {
26 layers: Vec<Fraction<FieldVec<P, A>>>,
27 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 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 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 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 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 pub const fn n_layers(&self) -> usize {
151 self.layers.len()
152 }
153
154 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 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 pub fn prove(
191 self,
192 claim: FracAddEvalClaim<F>,
193 channel: &mut impl IPProverChannel<F>,
194 ) -> FracAddEvalClaim<F> {
195 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 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 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 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 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 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 let (prover, sums) = FracAddCircuit::build(
294 k,
295 &alloc,
296 Fraction::new(witness_num.clone(), witness_den.clone()),
297 );
298
299 let eval_point = random_scalars::<P::Scalar>(&mut rng, n);
301
302 let sum_num_eval = evaluate(&sums.num, &eval_point);
304 let sum_den_eval = evaluate(&sums.den, &eval_point);
305 let claim = FracAddEvalClaim {
307 num_eval: sum_num_eval,
308 den_eval: sum_den_eval,
309 point: eval_point,
310 };
311
312 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
314 let prover_output = prover.prove(claim.clone(), &mut prover_transcript);
315
316 let mut verifier_transcript = prover_transcript.into_verifier();
318 let verifier_output = fracaddcheck::verify(k, claim, &mut verifier_transcript).unwrap();
319
320 assert_eq!(prover_output, verifier_output);
322
323 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 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 let (prover, sums) = FracAddCircuit::build(
350 k,
351 &alloc,
352 Fraction::new(witness_num.clone(), witness_den.clone()),
353 );
354
355 for (j, layer) in prover.layers.iter().chain(iter::once(&sums)).enumerate() {
357 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 #[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}