1use std::iter;
4
5use binius_compute::{Allocator, VecLike};
6use binius_field::{Field, PackedField};
7use binius_ip::{mlecheck, prodcheck::MultilinearEvalClaim};
8use binius_math::{
9 FieldBuffer, FieldVec,
10 line::extrapolate_line,
11 multilinear::eq::{eq_ind_partial_eval, eq_one_var},
12};
13use binius_utils::rayon::{
14 prelude::*,
15 task_size::{IndexedParallelIteratorExt, WorkPerItem},
16};
17use itertools::izip;
18
19use crate::{
20 channel::IPProverChannel,
21 sumcheck::{
22 ProveSingleOutput, bivariate_product_mle, common::MleCheckProver, prove_single_mlecheck,
23 },
24};
25
26pub mod one_pad_mle;
27
28use one_pad_mle::OnePadMleCheckProver;
29
30pub struct ProdcheckProver<'a, A: Allocator, P: PackedField> {
35 layers: Vec<FieldVec<P, A>>,
39 alloc: &'a A,
41}
42
43impl<A: Allocator, P: PackedField> Clone for ProdcheckProver<'_, A, P>
47where
48 A::Vec<P>: Clone,
49{
50 fn clone(&self) -> Self {
51 Self {
52 layers: self.layers.clone(),
53 alloc: self.alloc,
54 }
55 }
56}
57
58impl<'a, A, F, P> ProdcheckProver<'a, A, P>
59where
60 A: Allocator,
61 F: Field,
62 P: PackedField<Scalar = F>,
63{
64 pub fn new(k: usize, alloc: &'a A, witness: FieldVec<P, A>) -> (Self, FieldVec<P, A>) {
77 assert!(witness.log_len() >= k); let mut layers = Vec::with_capacity(k + 1);
80 layers.push(witness);
81
82 for _ in 0..k {
83 let prev_layer = layers.last().expect("layers is non-empty");
84 let next_log_len = prev_layer.log_len() - 1;
85 let (half_0, half_1) = prev_layer.split_half();
86
87 let out_len = half_0.as_ref().len();
91 let mut next_data = alloc.alloc::<P>(out_len);
92 (next_data.spare_capacity_mut(), half_0.as_ref(), half_1.as_ref())
93 .into_par_iter()
94 .with_min_task(WorkPerItem::FieldMuls)
95 .for_each(|(out, &v0, &v1)| {
96 out.write(v0 * v1);
97 });
98 assert!(
106 next_data.capacity() - next_data.len() >= out_len,
107 "the allocated buffer must hold every claimed slot"
108 );
109 assert_eq!(
110 half_1.as_ref().len(),
111 out_len,
112 "the two sibling halves must hold exactly one word per claimed slot"
113 );
114 unsafe {
119 next_data.set_len(out_len);
120 }
121 let next_layer = FieldBuffer::new(next_log_len, next_data);
122
123 layers.push(next_layer);
124 }
125
126 let products = layers.pop().expect("layers has k+1 elements");
127 (Self { layers, alloc }, products)
128 }
129
130 pub const fn n_layers(&self) -> usize {
132 self.layers.len()
133 }
134
135 pub fn pop_layer(&mut self) -> FieldVec<P, A> {
140 self.layers
141 .pop()
142 .expect("precondition: layers is non-empty")
143 }
144
145 pub fn layer_prover(
151 mut self,
152 claim: MultilinearEvalClaim<F>,
153 ) -> (impl MleCheckProver<F> + 'a, Option<Self>) {
154 let alloc = self.alloc;
155 let layer = self.layers.pop().expect("layers is non-empty");
156
157 let remaining = if self.layers.is_empty() {
158 None
159 } else {
160 Some(self)
161 };
162
163 let prover = bivariate_product_mle::new_split_half(alloc, layer, claim.point, claim.eval);
167
168 (prover, remaining)
169 }
170
171 pub fn prove(
183 self,
184 claim: MultilinearEvalClaim<F>,
185 channel: &mut impl IPProverChannel<F>,
186 ) -> MultilinearEvalClaim<F> {
187 let mut prover_opt = Some(self);
188 let mut claim = claim;
189
190 while let Some(prover) = prover_opt {
191 let (mle_prover, remaining) = prover.layer_prover(claim.clone());
192 prover_opt = remaining;
193
194 let ProveSingleOutput {
195 multilinear_evals,
196 challenges,
197 } = prove_single_mlecheck(mle_prover, channel);
198
199 let [eval_0, eval_1] = multilinear_evals
200 .try_into()
201 .expect("prover has two multilinears");
202
203 channel.send_many(&[eval_0, eval_1]);
204
205 let r = channel.sample();
206 let next_eval = extrapolate_line(eval_0, eval_1, r);
207
208 let mut next_point = challenges;
209 next_point.reverse();
210 next_point.push(r);
211
212 claim = MultilinearEvalClaim {
213 eval: next_eval,
214 point: next_point,
215 };
216 }
217
218 claim
219 }
220}
221
222pub struct BatchProveOutput<F> {
228 pub eval_point: Vec<F>,
230 pub evals: Vec<F>,
232}
233
234pub fn batch_prove<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
288 provers: Vec<ProdcheckProver<'a, A, P>>,
289 claimed_products: Vec<F>,
290 selector_point: Vec<F>,
291 content_point: Vec<F>,
292 channel: &mut impl IPProverChannel<F>,
293) -> BatchProveOutput<F> {
294 assert!(!provers.is_empty()); assert_eq!(claimed_products.len(), provers.len()); let n = provers.len();
298 let k = selector_point.len();
299 assert!(provers.len() <= (1 << k)); let n_layers = provers[0].n_layers();
302 assert!(n_layers >= 1); assert!(provers.iter().all(|p| p.n_layers() == n_layers)); let eval_point = [selector_point, content_point].concat();
309
310 let (provers, evals, eval_point) = (0..n_layers).fold(
311 (provers, claimed_products, eval_point),
312 |(provers, claimed_products, eval_point), _| {
313 batch_prove_layer(provers, claimed_products, &eval_point, k, channel)
314 },
315 );
316 debug_assert!(provers.is_empty(), "the final layer leaves no provers");
317 debug_assert_eq!(evals.len(), n);
318
319 BatchProveOutput { eval_point, evals }
320}
321
322#[allow(clippy::type_complexity)]
323fn batch_prove_layer<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
324 provers: Vec<ProdcheckProver<'a, A, P>>,
325 claimed_products: Vec<F>,
326 eval_point: &[F],
327 k: usize,
328 channel: &mut impl IPProverChannel<F>,
329) -> (Vec<ProdcheckProver<'a, A, P>>, Vec<F>, Vec<F>) {
330 let alloc = provers[0].alloc;
331 let inner_coords = &eval_point[k..];
332
333 let (layer_provers, next_provers): (Vec<_>, Vec<_>) = iter::zip(provers, claimed_products)
334 .map(|(prover, prod)| {
335 prover.layer_prover(MultilinearEvalClaim {
336 eval: prod,
337 point: inner_coords.to_vec(),
338 })
339 })
340 .unzip();
341
342 let (next_claimed_products, next_point) =
343 prove_layer_rounds::<A, F, P>(layer_provers, eval_point, k, alloc, channel);
344 let next_provers = next_provers.into_iter().flatten().collect();
345
346 (next_provers, next_claimed_products, next_point)
347}
348
349fn prove_layer_rounds<A: Allocator, F: Field, P: PackedField<Scalar = F>>(
360 mut layer_provers: Vec<impl MleCheckProver<F>>,
361 eval_point: &[F],
362 k: usize,
363 alloc: &A,
364 channel: &mut impl IPProverChannel<F>,
365) -> (Vec<F>, Vec<F>) {
366 let n = layer_provers.len();
367 let (outer_coords, inner_coords) = eval_point.split_at(k);
369
370 let eq_weights = eq_ind_partial_eval::<F>(outer_coords);
372
373 let mut challenges = Vec::with_capacity(eval_point.len());
375
376 for _round in 0..inner_coords.len() {
377 let coeffss = layer_provers
379 .iter_mut()
380 .map(|prover| {
381 let mut round_coeffs_vec = prover.execute();
382 round_coeffs_vec
383 .pop()
384 .expect("prodcheck layer provers have round_coeffs_vec.len() == 1")
385 })
386 .collect::<Vec<_>>();
387
388 let coeffs = iter::zip(coeffss, eq_weights.iter_scalars())
389 .map(|(coeffs, weight)| coeffs * weight)
390 .sum();
391
392 channel.send_many(mlecheck::RoundProof::truncate(coeffs).coeffs());
394
395 let challenge = channel.sample();
397 challenges.push(challenge);
398
399 for prover in layer_provers.iter_mut() {
400 prover.fold(challenge);
401 }
402 }
403
404 let (mut vals_0, mut vals_1): (Vec<F>, Vec<F>) = layer_provers
406 .into_iter()
407 .map(|prover| {
408 let evals = prover.finish();
409 let [e0, e1]: [F; 2] = evals
410 .try_into()
411 .expect("bivariate product prover has two multilinears");
412 (e0, e1)
413 })
414 .unzip();
415
416 vals_0.resize(1 << k, F::ZERO);
418 vals_1.resize(1 << k, F::ZERO);
419
420 let eval = izip!(&vals_0, &vals_1, eq_weights.as_ref())
422 .map(|(&v0, &v1, &eq_i)| v0 * v1 * eq_i)
423 .sum();
424
425 let outer_prover = bivariate_product_mle::new(
427 alloc,
428 [
429 FieldBuffer::<P, _>::from_values_in(alloc, &vals_0),
430 FieldBuffer::<P, _>::from_values_in(alloc, &vals_1),
431 ],
432 outer_coords.to_vec(),
433 eval,
434 );
435
436 let ProveSingleOutput {
437 multilinear_evals: outer_evals,
438 challenges: outer_challenges,
439 } = prove_single_mlecheck(outer_prover, channel);
440
441 challenges.extend(outer_challenges);
442
443 let [merged_eval_0, merged_eval_1]: [F; 2] =
444 outer_evals.try_into().expect("prover has two multilinears");
445
446 channel.send_many(&[merged_eval_0, merged_eval_1]);
448
449 let r = channel.sample();
450
451 let mut next_point = challenges;
452 next_point.reverse();
453 next_point.push(r);
454
455 let next_claimed_products = iter::zip(&vals_0[..n], &vals_1[..n])
457 .map(|(e0, e1)| extrapolate_line(*e0, *e1, r))
458 .collect();
459
460 (next_claimed_products, next_point)
461}
462
463pub struct BatchProveUnequalDepthsOutput<F, Prover> {
468 pub eval_point: Vec<F>,
470 pub provers: Vec<(F, Prover)>,
473}
474
475pub fn batch_prove_unequal_depths<'a, A, F, P, Channel>(
518 mut provers: Vec<ProdcheckProver<'a, A, P>>,
519 claimed_products: Vec<F>,
520 selector_point: Vec<F>,
521 channel: &mut Channel,
522) -> BatchProveUnequalDepthsOutput<F, impl MleCheckProver<F> + use<'a, A, F, P, Channel>>
523where
524 A: Allocator,
525 F: Field,
526 P: PackedField<Scalar = F>,
527 Channel: IPProverChannel<F>,
528{
529 assert!(!provers.is_empty()); assert_eq!(claimed_products.len(), provers.len()); let k = selector_point.len();
533 assert!(provers.len() <= (1 << k)); assert!(provers.iter().all(|prover| prover.n_layers() >= 1)); let alloc = provers[0].alloc;
537 let n_layers = provers
538 .iter()
539 .map(ProdcheckProver::n_layers)
540 .max()
541 .expect("provers is non-empty");
542 let pad_lens = provers
544 .iter()
545 .map(|prover| n_layers - prover.n_layers())
546 .collect::<Vec<_>>();
547
548 let products = claimed_products.clone();
551 let mut claims = claimed_products;
552 let mut eval_point = selector_point;
553
554 for _ in 0..n_layers - 1 {
557 let layer_provers =
558 layer_provers(&mut provers, &pad_lens, &products, &claims, &eval_point[k..]);
559 let (next_claims, next_point) =
560 prove_layer_rounds::<A, F, P>(layer_provers, &eval_point, k, alloc, channel);
561 claims = next_claims;
562 eval_point = next_point;
563 }
564
565 let provers = layer_provers(&mut provers, &pad_lens, &products, &claims, &eval_point[k..]);
566
567 BatchProveUnequalDepthsOutput {
568 eval_point,
569 provers: iter::zip(claims, provers).collect(),
570 }
571}
572
573fn layer_provers<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
575 provers: &mut [ProdcheckProver<'a, A, P>],
576 pad_lens: &[usize],
577 products: &[F],
578 claims: &[F],
579 node_point: &[F],
580) -> Vec<OnePadMleCheckProver<F, impl MleCheckProver<F> + use<'a, A, F, P>>> {
581 let node_len = node_point.len();
582
583 izip!(provers, pad_lens, products, claims)
584 .map(|(prover, &pad_len, &product, &claim)| {
585 let alloc = prover.alloc;
586 let layer = if node_len < pad_len {
590 FieldBuffer::from_values_in(alloc, &[product, F::ONE])
591 } else {
592 let layer = prover.pop_layer();
593 assert_eq!(
594 layer.log_len(),
595 node_len - pad_len + 1,
596 "precondition: the witness has exactly n_layers variables"
597 );
598 layer
599 };
600 one_pad_mle::new(alloc, layer, pad_len.min(node_len), node_point.to_vec(), claim)
601 })
602 .collect()
603}
604
605pub fn unpad_leaf_claim<F: Field>(
637 eval: F,
638 point: &[F],
639 n_pad_vars: usize,
640) -> MultilinearEvalClaim<F> {
641 assert!(point.len() >= n_pad_vars); let pad_eq = point[..n_pad_vars]
644 .iter()
645 .map(|&coord| eq_one_var(F::ZERO, coord))
646 .product::<F>();
647 assert!(pad_eq != F::ZERO, "a padding coordinate equals one");
648
649 MultilinearEvalClaim {
650 eval: F::ONE + (eval - F::ONE) * pad_eq.invert_or_zero(),
651 point: point[n_pad_vars..].to_vec(),
652 }
653}
654
655#[cfg(test)]
656mod tests {
657 use binius_field::{PackedField, field::FieldOps};
658 use binius_ip::prodcheck;
659 use binius_math::{
660 inner_product::inner_product,
661 multilinear::{eq::eq_ind_partial_eval, evaluate::evaluate},
662 test_utils::{Packed128b, random_field_buffer, random_scalars},
663 };
664 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
665 use binius_utils::checked_arithmetics::log2_ceil_usize;
666
667 type StdChallenger = HasherChallenger<sha2::Sha256>;
668 use binius_compute::GlobalAllocator;
669 use rand::prelude::*;
670
671 use super::*;
672
673 fn combine_batch_prove<F: Field, P: PackedField<Scalar = F>>(
677 output: BatchProveOutput<F>,
678 k: usize,
679 ) -> MultilinearEvalClaim<F> {
680 let BatchProveOutput { eval_point, evals } = output;
681 let eq_weights = eq_ind_partial_eval::<P>(&eval_point[..k]);
682 let final_eval =
683 inner_product(evals.iter().copied(), (0..evals.len()).map(|i| eq_weights.get(i)));
684
685 MultilinearEvalClaim {
686 eval: final_eval,
687 point: eval_point,
688 }
689 }
690
691 fn test_prodcheck_prove_verify_helper<P: PackedField>(n: usize, k: usize) {
692 let mut rng = StdRng::seed_from_u64(0);
693 let alloc = GlobalAllocator;
694
695 let witness = random_field_buffer::<P>(&mut rng, n + k);
697
698 let (prover, products) = ProdcheckProver::new(k, &alloc, witness.clone());
700
701 let eval_point = random_scalars::<P::Scalar>(&mut rng, n);
703
704 let products_eval = evaluate(&products, &eval_point);
706 let claim = MultilinearEvalClaim {
707 eval: products_eval,
708 point: eval_point,
709 };
710
711 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
713 let prover_output = prover.prove(claim.clone(), &mut prover_transcript);
714
715 let mut verifier_transcript = prover_transcript.into_verifier();
717 let verifier_output = prodcheck::verify(k, claim, &mut verifier_transcript).unwrap();
718
719 assert_eq!(prover_output, verifier_output);
721
722 let expected_eval = evaluate(&witness, &verifier_output.point);
724 assert_eq!(verifier_output.eval, expected_eval);
725 }
726
727 #[test]
728 fn test_prodcheck_prove_verify() {
729 test_prodcheck_prove_verify_helper::<Packed128b>(4, 3);
730 }
731
732 #[test]
733 fn test_prodcheck_full_prove_verify() {
734 test_prodcheck_prove_verify_helper::<Packed128b>(0, 4);
735 }
736
737 fn test_prodcheck_layer_computation_helper<P: PackedField>(n: usize, k: usize) {
738 let mut rng = StdRng::seed_from_u64(0);
739 let alloc = GlobalAllocator;
740
741 let witness = random_field_buffer::<P>(&mut rng, n + k);
743
744 let (_prover, products) = ProdcheckProver::new(k, &alloc, witness.clone());
746
747 let stride = 1 << n;
750 let num_terms = 1 << k;
751 for i in 0..(1 << n) {
752 let mut expected_product = P::Scalar::ONE;
753 for z in 0..num_terms {
754 expected_product *= witness.get(i + z * stride);
755 }
756 let actual = products.get(i);
757 assert_eq!(actual, expected_product, "Product mismatch at index {i}");
758 }
759 }
760
761 #[test]
762 fn test_prodcheck_layer_computation() {
763 test_prodcheck_layer_computation_helper::<Packed128b>(4, 3);
764 }
765
766 fn reference_layers<P: PackedField>(
770 k: usize,
771 alloc: &GlobalAllocator,
772 witness: FieldBuffer<P>,
773 ) -> Vec<FieldBuffer<P>> {
774 let mut layers = Vec::with_capacity(k + 1);
775 layers.push(witness);
776
777 for _ in 0..k {
778 let prev_layer = layers.last().expect("layers is non-empty");
779 let next_log_len = prev_layer.log_len() - 1;
780 let (half_0, half_1) = prev_layer.split_half();
781
782 let evals = half_0
784 .as_ref()
785 .iter()
786 .zip(half_1.as_ref())
787 .map(|(v0, v1)| *v0 * *v1)
788 .collect::<Vec<P>>();
789
790 let mut next_data = alloc.alloc::<P>(evals.len());
791 next_data.extend_from_slice(&evals);
792 layers.push(FieldBuffer::new(next_log_len, next_data));
793 }
794
795 layers
796 }
797
798 #[test]
799 fn layers_match_the_reference_word_for_word() {
800 let mut rng = StdRng::seed_from_u64(7);
801 let alloc = GlobalAllocator;
802
803 for log_len in 1..=6 {
814 for k in 1..=log_len {
815 let witness = random_field_buffer::<Packed128b>(&mut rng, log_len);
816
817 let (prover, products) = ProdcheckProver::new(k, &alloc, witness.clone());
818 let reference = reference_layers(k, &alloc, witness);
819
820 let built = prover.layers.iter().chain(iter::once(&products));
822
823 for (depth, (built, reference)) in built.zip(&reference).enumerate() {
824 assert_eq!(built.log_len(), reference.log_len(), "log_len at depth {depth}");
825 assert_eq!(
826 built.as_ref(),
827 reference.as_ref(),
828 "layer words at depth {depth}, witness log_len {log_len}, k {k}"
829 );
830 }
831 }
832 }
833 }
834
835 fn test_batch_prove_verify_helper<P: PackedField>(n_layers: usize, n_provers: usize) {
845 let mut rng = StdRng::seed_from_u64(42);
846 let alloc = GlobalAllocator;
847
848 let log_n_provers = log2_ceil_usize(n_provers);
849
850 let witnesses: Vec<FieldBuffer<P>> = (0..n_provers)
852 .map(|_| random_field_buffer::<P>(&mut rng, n_layers))
853 .collect();
854
855 let (provers, individual_products): (Vec<_>, Vec<_>) = witnesses
857 .iter()
858 .map(|witness| ProdcheckProver::new(n_layers, &alloc, witness.clone()))
859 .unzip();
860
861 let claimed_products: Vec<P::Scalar> = individual_products
863 .iter()
864 .map(|products| {
865 assert_eq!(products.log_len(), 0);
866 products.get(0)
867 })
868 .collect();
869
870 let selector_challenge = random_scalars::<P::Scalar>(&mut rng, log_n_provers);
872
873 let eq_weights = eq_ind_partial_eval::<P>(&selector_challenge);
875 let combined_eval = inner_product(
876 claimed_products.iter().copied(),
877 (0..n_provers).map(|i| eq_weights.get(i)),
878 );
879
880 let claim = MultilinearEvalClaim {
882 eval: combined_eval,
883 point: selector_challenge.clone(),
884 };
885
886 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
888 let batch_output = batch_prove(
889 provers,
890 claimed_products,
891 selector_challenge,
892 Vec::new(),
893 &mut prover_transcript,
894 );
895 assert_eq!(batch_output.evals.len(), n_provers);
897 let prover_output = combine_batch_prove::<_, P>(batch_output, log_n_provers);
898
899 let mut verifier_transcript = prover_transcript.into_verifier();
901 let verifier_output = prodcheck::verify(n_layers, claim, &mut verifier_transcript).unwrap();
902
903 assert_eq!(prover_output, verifier_output);
905
906 let final_point = &verifier_output.point;
908 assert_eq!(final_point.len(), log_n_provers + n_layers);
909
910 let selector_challenges = &final_point[..log_n_provers];
911 let content_challenges = &final_point[log_n_provers..];
912
913 let selector_weights = eq_ind_partial_eval::<P>(selector_challenges);
914
915 let expected_eval: P::Scalar = inner_product(
916 (0..n_provers).map(|i| evaluate(&witnesses[i], content_challenges)),
917 (0..n_provers).map(|i| selector_weights.get(i)),
918 );
919
920 assert_eq!(
921 verifier_output.eval, expected_eval,
922 "Final evaluation should match batch witness interpolation"
923 );
924 }
925
926 #[test]
927 fn test_batch_prove_power_of_two_provers() {
928 test_batch_prove_verify_helper::<Packed128b>(3, 4);
930 }
931
932 #[test]
933 fn test_batch_prove_non_power_of_two_provers() {
934 test_batch_prove_verify_helper::<Packed128b>(4, 3);
936 }
937
938 #[test]
939 fn test_batch_prove_single_prover() {
940 test_batch_prove_verify_helper::<Packed128b>(5, 1);
942 }
943
944 #[test]
945 fn test_batch_prove_single_layer() {
946 test_batch_prove_verify_helper::<Packed128b>(1, 4);
949 }
950
951 fn test_batch_prove_with_content_helper<P: PackedField>(
960 n_layers: usize,
961 n_provers: usize,
962 content_len: usize,
963 ) {
964 let mut rng = StdRng::seed_from_u64(7);
965 let alloc = GlobalAllocator;
966
967 let log_n_provers = log2_ceil_usize(n_provers);
968
969 let witnesses: Vec<FieldBuffer<P>> = (0..n_provers)
971 .map(|_| random_field_buffer::<P>(&mut rng, content_len + n_layers))
972 .collect();
973
974 let (provers, individual_products): (Vec<_>, Vec<_>) = witnesses
975 .iter()
976 .map(|witness| ProdcheckProver::new(n_layers, &alloc, witness.clone()))
977 .unzip();
978
979 let content_point = random_scalars::<P::Scalar>(&mut rng, content_len);
981 let claimed_products: Vec<P::Scalar> = individual_products
982 .iter()
983 .map(|products| {
984 assert_eq!(products.log_len(), content_len);
985 evaluate(products, &content_point)
986 })
987 .collect();
988
989 let selector_challenge = random_scalars::<P::Scalar>(&mut rng, log_n_provers);
992 let eq_weights = eq_ind_partial_eval::<P>(&selector_challenge);
993 let combined_eval = inner_product(
994 claimed_products.iter().copied(),
995 (0..n_provers).map(|i| eq_weights.get(i)),
996 );
997
998 let claim = MultilinearEvalClaim {
999 eval: combined_eval,
1000 point: [selector_challenge.clone(), content_point.clone()].concat(),
1001 };
1002
1003 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
1005 let batch_output = batch_prove(
1006 provers,
1007 claimed_products,
1008 selector_challenge,
1009 content_point,
1010 &mut prover_transcript,
1011 );
1012 assert_eq!(batch_output.evals.len(), n_provers);
1013 let prover_output = combine_batch_prove::<_, P>(batch_output, log_n_provers);
1014
1015 let mut verifier_transcript = prover_transcript.into_verifier();
1017 let verifier_output = prodcheck::verify(n_layers, claim, &mut verifier_transcript).unwrap();
1018
1019 assert_eq!(prover_output, verifier_output);
1020
1021 let final_point = &verifier_output.point;
1026 assert_eq!(final_point.len(), log_n_provers + n_layers + content_len);
1027
1028 let selector_challenges = &final_point[..log_n_provers];
1029 let witness_challenges = &final_point[log_n_provers..];
1030
1031 let selector_weights = eq_ind_partial_eval::<P>(selector_challenges);
1032
1033 let expected_eval: P::Scalar = inner_product(
1034 (0..n_provers).map(|i| evaluate(&witnesses[i], witness_challenges)),
1035 (0..n_provers).map(|i| selector_weights.get(i)),
1036 );
1037
1038 assert_eq!(
1039 verifier_output.eval, expected_eval,
1040 "Final evaluation should match batch witness interpolation"
1041 );
1042 }
1043
1044 #[test]
1045 fn test_batch_prove_with_content() {
1046 test_batch_prove_with_content_helper::<Packed128b>(4, 3, 2);
1048 }
1049
1050 #[allow(clippy::type_complexity)]
1054 fn unequal_depth_provers<'a, P: PackedField>(
1055 rng: &mut impl Rng,
1056 alloc: &'a GlobalAllocator,
1057 depths: &[usize],
1058 ) -> (Vec<FieldBuffer<P>>, Vec<ProdcheckProver<'a, GlobalAllocator, P>>, Vec<P::Scalar>) {
1059 itertools::multiunzip(depths.iter().map(|&depth| {
1060 let witness = random_field_buffer::<P>(&mut *rng, depth);
1061 let (prover, products) = ProdcheckProver::new(depth, alloc, witness.clone());
1062 assert_eq!(products.log_len(), 0);
1063 (witness, prover, products.get(0))
1064 }))
1065 }
1066
1067 fn combine_claims<P: PackedField>(
1069 claims: &[P::Scalar],
1070 selector_point: &[P::Scalar],
1071 ) -> P::Scalar {
1072 let eq_weights = eq_ind_partial_eval::<P>(selector_point);
1073 inner_product(claims.iter().copied(), (0..claims.len()).map(|i| eq_weights.get(i)))
1074 }
1075
1076 fn test_unequal_depths_helper<P: PackedField>(depths: &[usize]) {
1079 let mut rng = StdRng::seed_from_u64(11);
1080 let alloc = GlobalAllocator;
1081
1082 let k = log2_ceil_usize(depths.len());
1083 let n_layers = *depths.iter().max().expect("depths is non-empty");
1084
1085 let (witnesses, provers, claimed_products) =
1086 unequal_depth_provers::<P>(&mut rng, &alloc, depths);
1087
1088 let selector_point = random_scalars::<P::Scalar>(&mut rng, k);
1090 let claim = MultilinearEvalClaim {
1091 eval: combine_claims::<P>(&claimed_products, &selector_point),
1092 point: selector_point.clone(),
1093 };
1094
1095 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
1096 let BatchProveUnequalDepthsOutput {
1097 eval_point,
1098 provers,
1099 } = batch_prove_unequal_depths(
1100 provers,
1101 claimed_products,
1102 selector_point,
1103 &mut prover_transcript,
1104 );
1105
1106 let (_claims, provers): (Vec<_>, Vec<_>) = provers.into_iter().unzip();
1108 let (evals, eval_point) =
1109 prove_layer_rounds::<_, _, P>(provers, &eval_point, k, &alloc, &mut prover_transcript);
1110 assert_eq!(evals.len(), depths.len());
1111
1112 let mut verifier_transcript = prover_transcript.into_verifier();
1114 let verifier_output = prodcheck::verify(n_layers, claim, &mut verifier_transcript).unwrap();
1115
1116 assert_eq!(verifier_output.point, eval_point);
1117 assert_eq!(verifier_output.eval, combine_claims::<P>(&evals, &eval_point[..k]));
1118
1119 for (i, (&depth, witness)) in iter::zip(depths, &witnesses).enumerate() {
1122 let leaf = unpad_leaf_claim(evals[i], &eval_point[k..], n_layers - depth);
1123 assert_eq!(leaf.point.len(), depth);
1124 assert_eq!(leaf.eval, evaluate(witness, &leaf.point), "tree {i}");
1125 }
1126 }
1127
1128 #[test]
1129 fn test_unequal_depths_mixed() {
1130 test_unequal_depths_helper::<Packed128b>(&[2, 4, 5]);
1131 }
1132
1133 #[test]
1134 fn test_unequal_depths_single_prover() {
1135 test_unequal_depths_helper::<Packed128b>(&[3]);
1136 }
1137
1138 #[test]
1139 fn test_unequal_depths_power_of_two_provers() {
1140 test_unequal_depths_helper::<Packed128b>(&[1, 2, 5, 5]);
1142 }
1143
1144 #[test]
1145 fn test_unequal_depths_all_minimal() {
1146 test_unequal_depths_helper::<Packed128b>(&[1, 1, 1]);
1148 }
1149
1150 #[test]
1151 fn test_unequal_depths_maximal_padding() {
1152 test_unequal_depths_helper::<Packed128b>(&[1, 6]);
1154 }
1155
1156 #[test]
1159 fn test_unequal_depths_matches_batch_prove_at_equal_depths() {
1160 type P = Packed128b;
1161 type F = <P as FieldOps>::Scalar;
1162
1163 let depths = [4; 3];
1164 let k = log2_ceil_usize(depths.len());
1165 let alloc = GlobalAllocator;
1166
1167 let mut rng = StdRng::seed_from_u64(23);
1168 let selector_point = random_scalars::<F>(&mut rng, k);
1169 let prover_seed = 24;
1171
1172 let unequal_proof = {
1173 let mut rng = StdRng::seed_from_u64(prover_seed);
1174 let (_, provers, claimed_products) =
1175 unequal_depth_provers::<P>(&mut rng, &alloc, &depths);
1176
1177 let mut transcript = ProverTranscript::new(StdChallenger::default());
1178 let BatchProveUnequalDepthsOutput {
1179 eval_point,
1180 provers,
1181 } = batch_prove_unequal_depths(
1182 provers,
1183 claimed_products,
1184 selector_point.clone(),
1185 &mut transcript,
1186 );
1187 let (_claims, provers): (Vec<_>, Vec<_>) = provers.into_iter().unzip();
1188 prove_layer_rounds::<_, _, P>(provers, &eval_point, k, &alloc, &mut transcript);
1189 transcript.finalize()
1190 };
1191
1192 let equal_proof = {
1193 let mut rng = StdRng::seed_from_u64(prover_seed);
1194 let (_, provers, claimed_products) =
1195 unequal_depth_provers::<P>(&mut rng, &alloc, &depths);
1196
1197 let mut transcript = ProverTranscript::new(StdChallenger::default());
1198 batch_prove(provers, claimed_products, selector_point, Vec::new(), &mut transcript);
1199 transcript.finalize()
1200 };
1201
1202 assert_eq!(unequal_proof, equal_proof);
1203 }
1204}