Skip to main content

binius_ip_prover/sumcheck/
bivariate_product_evaluator.rs

1// Copyright 2026 The Binius Developers
2
3use binius_compute::Allocator;
4use binius_field::{Field, PackedField, WideMul};
5use binius_ip::sumcheck::RoundCoeffs;
6use binius_math::FieldVec;
7use itertools::izip;
8
9use super::{
10	mle_store::{ColId, EvaluationChunk, MleStore, RoundContext},
11	round_evals::RoundEvals,
12	round_evaluator::{SharedSumcheckProver, SumcheckRoundEvaluator},
13};
14
15/// Sumcheck round evaluator for a composite defined as the product of two store columns.
16///
17/// This is the store-backed counterpart of the bivariate product sumcheck prover: it proves the
18/// plain (non-eq-weighted) sum claim of the product over the hypercube, emitting regular
19/// sumcheck round polynomials.
20pub struct BivariateProductEvaluator {
21	cols: [ColId; 2],
22}
23
24impl BivariateProductEvaluator {
25	/// Creates an evaluator for the product of two store columns.
26	///
27	/// The claimed sum is held by the driving [`SharedSumcheckProver`], not the evaluator.
28	pub const fn new(cols: [ColId; 2]) -> Self {
29		Self { cols }
30	}
31}
32
33/// Builds a [`SumcheckProver`](crate::sumcheck::common::SumcheckProver) for the plain hypercube sum
34/// of the product of two multilinears.
35///
36/// This is the store-backed replacement for the former bespoke bivariate-product prover: it loads
37/// the two columns into a fresh [`MleStore`] and drives a single [`BivariateProductEvaluator`]. The
38/// multilinears must have the same number of variables; the returned prover's `finish` emits their
39/// two evaluations in the given order.
40pub fn bivariate_product_prover<'alloc, A: Allocator, F: Field, P: PackedField<Scalar = F>>(
41	alloc: &'alloc A,
42	multilinears: [FieldVec<P, A>; 2],
43	sum: F,
44) -> SharedSumcheckProver<'alloc, A, P, BivariateProductEvaluator> {
45	assert_eq!(
46		multilinears[0].log_len(),
47		multilinears[1].log_len(),
48		"multilinears must have equal number of variables"
49	);
50
51	let mut store = MleStore::new(multilinears[0].log_len(), alloc);
52	let cols = multilinears.map(|col| store.push_owned(col));
53	SharedSumcheckProver::new(store, [(sum, BivariateProductEvaluator::new(cols))])
54}
55
56impl<F: Field, P: PackedField<Scalar = F>> SumcheckRoundEvaluator<F, P>
57	for BivariateProductEvaluator
58{
59	fn degree(&self) -> usize {
60		// Product of two multilinears: two sampled evaluations, `y_1` and `y_inf`.
61		2
62	}
63
64	fn accumulate(&self, chunk: &EvaluationChunk<'_, P>, accum: &mut [<P as WideMul>::Output]) {
65		let a = chunk.col(self.cols[0]);
66		let b = chunk.col(self.cols[1]);
67
68		// Accumulate F(1) and F(∞) where F = ∑_{v ∈ B} A(v || X) B(v || X). The low half is the
69		// v-prefix at x=0, the high half at x=1.
70		//
71		// The per-point products are accumulated in unreduced (wide) form and reduced a single
72		// time in interpolate, amortizing the GF(2^128) reduction over the whole sum.
73		let mut y_1 = <P as WideMul>::Output::default();
74		let mut y_inf = <P as WideMul>::Output::default();
75		for (&a_0_i, &a_1_i, &b_0_i, &b_1_i) in
76			izip!(a.lo.as_ref(), a.hi.as_ref(), b.lo.as_ref(), b.hi.as_ref())
77		{
78			// Evaluate M(∞) = M(0) + M(1)
79			let a_inf_i = a_0_i + a_1_i;
80			let b_inf_i = b_0_i + b_1_i;
81
82			y_1 += P::wide_mul(a_1_i, b_1_i);
83			y_inf += P::wide_mul(a_inf_i, b_inf_i);
84		}
85
86		RoundEvals([y_1, y_inf]).add_to(accum);
87	}
88
89	fn interpolate(&self, ctx: &RoundContext<'_, P>, accum: &[P], claim: F) -> RoundCoeffs<F> {
90		// The store has not yet folded this round, so its remaining-variable count is this round's.
91		let n_vars_remaining = ctx.n_vars();
92		assert!(n_vars_remaining > 0);
93
94		// `claim` is this round's sum.
95		// The evaluation at 0, which was never sampled, is recovered from it.
96		RoundEvals::<P, 2>::from_slots(accum)
97			.sum_scalars(n_vars_remaining)
98			.interpolate(claim)
99	}
100}
101
102#[cfg(test)]
103mod tests {
104	use binius_compute::GlobalAllocator;
105	use binius_field::arch::{OptimalB128, OptimalPackedB128};
106	use binius_ip::sumcheck::verify;
107	use binius_math::{
108		inner_product::inner_product_par, multilinear::evaluate::evaluate,
109		test_utils::random_field_buffer,
110	};
111	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
112	use rand::prelude::*;
113
114	use super::*;
115	use crate::sumcheck::prove::prove_single;
116
117	type StdChallenger = HasherChallenger<sha2::Sha256>;
118
119	// Proving the product sum of two multilinears via the shared store, then verifying, recovers
120	// the two multilinear evaluations at the challenge point and their product as the reduced
121	// eval.
122	#[test]
123	fn test_bivariate_product_sumcheck() {
124		type F = OptimalB128;
125		type P = OptimalPackedB128;
126
127		let n_vars = 8;
128		let mut rng = StdRng::seed_from_u64(0);
129		let alloc = GlobalAllocator;
130
131		let multilinear_a = random_field_buffer::<P>(&mut rng, n_vars);
132		let multilinear_b = random_field_buffer::<P>(&mut rng, n_vars);
133		let expected_sum = inner_product_par(&multilinear_a, &multilinear_b);
134
135		let prover = bivariate_product_prover(
136			&alloc,
137			[multilinear_a.clone(), multilinear_b.clone()],
138			expected_sum,
139		);
140
141		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
142		let output = prove_single(prover, &mut prover_transcript);
143		prover_transcript
144			.message()
145			.write_slice(&output.multilinear_evals);
146
147		let mut verifier_transcript = prover_transcript.into_verifier();
148		let sumcheck_output = verify(n_vars, 2, expected_sum, &mut verifier_transcript).unwrap();
149		let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(2).unwrap();
150
151		assert_eq!(
152			multilinear_evals[0] * multilinear_evals[1],
153			sumcheck_output.eval,
154			"product of the multilinear evaluations should equal the reduced evaluation"
155		);
156
157		// The prover binds variables high-to-low; `evaluate` expects them low-to-high.
158		let mut eval_point = sumcheck_output.challenges.clone();
159		eval_point.reverse();
160		assert_eq!(evaluate(&multilinear_a, &eval_point), multilinear_evals[0]);
161		assert_eq!(evaluate(&multilinear_b, &eval_point), multilinear_evals[1]);
162		assert_eq!(output.challenges, sumcheck_output.challenges);
163	}
164}