1use 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
27pub struct BatchProveOutput<F> {
33 pub eval_point: Vec<F>,
35 pub fractions: Vec<Fraction<F>>,
37}
38
39fn 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 let (outer_coords, inner_coords) = eval_point.split_at(k);
68
69 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
101fn 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 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 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
152fn 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 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
204fn 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 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
267fn 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 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
303pub 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()); assert_eq!(claimed_fractions.len(), provers.len()); let k = selector_point.len();
355 assert!(provers.len() <= (1 << k)); 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 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 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 type Witness<P> = Fraction<FieldBuffer<P>>;
406
407 #[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 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 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 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 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 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 test_unequal_depths_helper::<Packed128b>(&[1, 2, 5, 5], 11);
514 }
515
516 #[test]
517 fn test_unequal_depths_all_minimal() {
518 test_unequal_depths_helper::<Packed128b>(&[1, 1, 1], 11);
520 }
521
522 #[test]
523 fn test_unequal_depths_zero_depth_tree() {
524 test_unequal_depths_helper::<Packed128b>(&[0, 3], 11);
526 }
527
528 #[test]
529 fn test_unequal_depths_maximal_padding() {
530 test_unequal_depths_helper::<Packed128b>(&[1, 6], 11);
532 }
533
534 #[test]
535 fn test_unequal_depths_equal_depths() {
536 test_unequal_depths_helper::<Packed128b>(&[4, 4, 4], 11);
538 }
539
540 proptest! {
541 #![proptest_config(ProptestConfig::with_cases(64))]
543
544 #[test]
549 fn unequal_depths_round_trip(
550 seed in any::<u64>(),
551 depths in prop::collection::vec(0usize..=6, 1..=5),
552 ) {
553 prop_assume!(depths.iter().any(|&depth| depth > 0));
555
556 test_unequal_depths_helper::<Packed128b>(&depths, seed);
557 }
558 }
559}