Skip to main content

binius_ip_prover/sumcheck/
frac_add_mle.rs

1// Copyright 2025-2026 The Binius Developers
2
3use binius_compute::Allocator;
4use binius_field::{Field, PackedField};
5use binius_math::FieldVec;
6
7use crate::sumcheck::{
8	mle_store::{ColId, MleStore},
9	quadratic_mle_evaluator::QuadraticMleEvaluator,
10	round_evaluator::{MleCheckRoundEvaluator, SharedMleCheckProver},
11};
12
13/// The store-based MLE-check prover for one fractional-addition layer.
14///
15/// It owns its four half-columns, so it is self-contained: a caller can drive it, batch it, or
16/// extend its store with more columns and evaluators.
17pub type LayerProver<'a, A, F, P> =
18	SharedMleCheckProver<'a, A, F, P, Box<dyn MleCheckRoundEvaluator<F, P> + 'a>>;
19
20/// Creates the [`LayerProver`] reducing one fractional-addition layer, sharing two allocations.
21///
22/// `num` and `den` each have one more variable than `eval_point`:
23/// - their low halves fix the highest variable to 0,
24/// - their high halves fix it to 1.
25///
26/// The reduction proves the fractional addition of those halves. Both halves of each buffer live
27/// inside the one allocation, so separating them costs no copy.
28///
29/// # Arguments
30///
31/// * `num`, `den` - The layer's numerator and denominator buffers.
32/// * `eval_point` - The shared point at which both claims are taken.
33/// * `claims` - The layer's claimed numerator and denominator evaluations, in that order.
34///
35/// # Preconditions
36/// * `num.log_len() == den.log_len()`
37/// * `num.log_len() == eval_point.len() + 1`
38pub fn new_split_half<'alloc, A, F, P>(
39	alloc: &'alloc A,
40	num: FieldVec<P, A>,
41	den: FieldVec<P, A>,
42	eval_point: Vec<F>,
43	claims: [F; 2],
44) -> LayerProver<'alloc, A, F, P>
45where
46	A: Allocator,
47	F: Field,
48	P: PackedField<Scalar = F>,
49{
50	assert_eq!(
51		num.log_len(),
52		den.log_len(),
53		"precondition: numerator and denominator have equal length"
54	);
55	// The store checks that each buffer has exactly one more variable than itself.
56	let mut store = MleStore::new(num.log_len() - 1, alloc);
57	let [num_0, num_1] = store.push_split_half(num);
58	let [den_0, den_1] = store.push_split_half(den);
59	let cols = [num_0, num_1, den_0, den_1];
60	let (num_evaluator, den_evaluator) = evaluators::<F, P>(cols);
61
62	let [num_claim, den_claim] = claims;
63	let claims_with_evaluators: [(F, Box<dyn MleCheckRoundEvaluator<F, P> + 'alloc>); 2] = [
64		(num_claim, Box::new(num_evaluator)),
65		(den_claim, Box::new(den_evaluator)),
66	];
67	SharedMleCheckProver::new(store, claims_with_evaluators, eval_point)
68}
69
70/// Creates the round evaluators for the fractional-addition claims required in logUp*.
71///
72/// The columns are `[num_a, num_b, den_a, den_b]`: the numerators and denominators of the two
73/// fraction collections being added, split as either half of one layer buffer. The two claims —
74/// one evaluator each — are the fractional-addition numerator `num_a * den_b + num_b * den_a` over
75/// all four columns and the denominator `den_a * den_b` over `[den_a, den_b]`, both weighted by the
76/// equality indicator at the shared evaluation point. The driving [`SharedMleCheckProver`] owns
77/// that point's eq tracker and holds the two claimed evaluations, so the evaluators carry neither.
78///
79/// [`SharedMleCheckProver`]: crate::sumcheck::round_evaluator::SharedMleCheckProver
80pub fn evaluators<F, P>(
81	cols: [ColId; 4],
82) -> (impl MleCheckRoundEvaluator<F, P> + 'static, impl MleCheckRoundEvaluator<F, P> + 'static)
83where
84	F: Field,
85	P: PackedField<Scalar = F>,
86{
87	let [num_a, num_b, den_a, den_b] = cols;
88	// The fractional addition formulas are purely quadratic, so each infinity composition matches
89	// its regular composition.
90	let numerator = QuadraticMleEvaluator::new(
91		[num_a, num_b, den_a, den_b],
92		|[num_a, num_b, den_a, den_b]: [P; 4]| num_a * den_b + num_b * den_a,
93		|[num_a, num_b, den_a, den_b]: [P; 4]| num_a * den_b + num_b * den_a,
94	);
95	let denominator = QuadraticMleEvaluator::new(
96		[den_a, den_b],
97		|[den_a, den_b]: [P; 2]| den_a * den_b,
98		|[den_a, den_b]: [P; 2]| den_a * den_b,
99	);
100	(numerator, denominator)
101}
102
103#[cfg(test)]
104mod tests {
105	use binius_field::arch::{OptimalB128, OptimalPackedB128};
106	use binius_ip::sumcheck::batch_verify;
107	use binius_math::{
108		FieldBuffer,
109		multilinear::{eq::eq_ind, evaluate::evaluate},
110		test_utils::{random_field_buffer, random_scalars},
111	};
112	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
113
114	type StdChallenger = HasherChallenger<sha2::Sha256>;
115	use binius_compute::GlobalAllocator;
116	use itertools::{Itertools, izip};
117	use rand::prelude::*;
118
119	use super::*;
120	use crate::sumcheck::{
121		MleToSumCheckEvaluator,
122		batch::batch_prove,
123		common::SumcheckProver,
124		mle_store::MleStore,
125		round_evaluator::{SharedSumcheckProver, SumcheckRoundEvaluator},
126	};
127
128	fn test_frac_add_sumcheck_prove_verify<F, P>(
129		prover: impl SumcheckProver<F>,
130		eval_claims: [F; 2],
131		eval_point: &[F],
132		num_a: &FieldBuffer<P>,
133		num_b: &FieldBuffer<P>,
134		den_a: &FieldBuffer<P>,
135		den_b: &FieldBuffer<P>,
136	) where
137		F: Field,
138		P: PackedField<Scalar = F>,
139	{
140		let n_vars = prover.n_vars();
141		// Run the proving protocol
142		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
143		let output = batch_prove(vec![prover], &mut prover_transcript);
144
145		assert_eq!(output.multilinear_evals.len(), 1);
146		let prover_evals = output.multilinear_evals[0].clone();
147
148		// Write the multilinear evaluations to the transcript
149		prover_transcript
150			.message()
151			.write_scalar_slice(&prover_evals);
152
153		// Convert to verifier transcript and run verification
154		let mut verifier_transcript = prover_transcript.into_verifier();
155		let sumcheck_output =
156		// Degree 3 because quadratic prime polynomials are multiplied by a linear eq term.
157		batch_verify(n_vars, 3, &eval_claims, &mut verifier_transcript).unwrap();
158
159		// The prover binds variables from high to low, but evaluate expects them from low to high
160		let mut reduced_eval_point = sumcheck_output.challenges.clone();
161		reduced_eval_point.reverse();
162
163		// Read the multilinear evaluations from the transcript
164		let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(4).unwrap();
165
166		// Evaluate the equality indicator
167		let eq_ind_eval = eq_ind(eval_point, &reduced_eval_point);
168
169		// Check that the original multilinears evaluate to the claimed values at the challenge
170		// point
171		let eval_num_a = evaluate(num_a, &reduced_eval_point);
172		let eval_den_a = evaluate(den_a, &reduced_eval_point);
173		let eval_num_b = evaluate(num_b, &reduced_eval_point);
174		let eval_den_b = evaluate(den_b, &reduced_eval_point);
175
176		assert_eq!(
177			eval_num_a, multilinear_evals[0],
178			"Numerator A should evaluate to the first claimed evaluation"
179		);
180
181		assert_eq!(
182			eval_num_b, multilinear_evals[1],
183			"Numerator B should evaluate to the second claimed evaluation"
184		);
185		assert_eq!(
186			eval_den_a, multilinear_evals[2],
187			"Denominator A should evaluate to the third claimed evaluation"
188		);
189
190		assert_eq!(
191			eval_den_b, multilinear_evals[3],
192			"Denominator B should evaluate to the fourth claimed evaluation"
193		);
194
195		// Check that the batched evaluation matches the sumcheck output
196		// Sumcheck wraps the prime polynomial with an eq factor, so include eq_ind_eval here.
197		let numerator_eval = (eval_num_a * eval_den_b + eval_num_b * eval_den_a) * eq_ind_eval;
198		let denominator_eval = (eval_den_a * eval_den_b) * eq_ind_eval;
199		let batched_eval = numerator_eval + denominator_eval * sumcheck_output.batch_coeff;
200
201		assert_eq!(
202			batched_eval, sumcheck_output.eval,
203			"Batched evaluation should equal the reduced evaluation"
204		);
205
206		assert_eq!(
207			output.challenges, sumcheck_output.challenges,
208			"Prover and verifier challenges should match"
209		);
210	}
211
212	#[test]
213	fn test_frac_add_sumcheck() {
214		type F = OptimalB128;
215		type P = OptimalPackedB128;
216
217		let n_vars = 8;
218		let mut rng = StdRng::seed_from_u64(0);
219		let alloc = GlobalAllocator;
220
221		let num_a = random_field_buffer::<P>(&mut rng, n_vars);
222		let num_b = random_field_buffer::<P>(&mut rng, n_vars);
223		let den_a = random_field_buffer::<P>(&mut rng, n_vars);
224		let den_b = random_field_buffer::<P>(&mut rng, n_vars);
225
226		let numerator_values =
227			izip!(num_a.as_ref(), den_a.as_ref(), num_b.as_ref(), den_b.as_ref())
228				.map(|(&num_a, &den_a, &num_b, &den_b)| num_a * den_b + num_b * den_a)
229				.collect_vec();
230
231		let denominator_values = izip!(den_a.as_ref(), den_b.as_ref())
232			.map(|(&den_a, &den_b)| den_a * den_b)
233			.collect_vec();
234
235		let numerator_buffer = FieldBuffer::new(n_vars, numerator_values);
236		let denominator_buffer = FieldBuffer::new(n_vars, denominator_values);
237
238		let eval_point = random_scalars::<F>(&mut rng, n_vars);
239		// Claims are at the original eval_point; verifier handles challenge ordering separately.
240		let eval_claims = [
241			evaluate(&numerator_buffer, &eval_point),
242			evaluate(&denominator_buffer, &eval_point),
243		];
244
245		let mut store = MleStore::new(n_vars, &alloc);
246		let cols = [num_a.clone(), num_b.clone(), den_a.clone(), den_b.clone()]
247			.map(|col| store.push_owned(col));
248		let (num_evaluator, den_evaluator) = evaluators(cols);
249
250		// Register the shared point's eq tracker for the sumcheck wrappers (a plain sumcheck prover
251		// knows nothing of eq indicators, so the wrappers hold the tracker themselves).
252		let eq_tracker = store.register_eq_tracker(&eval_point);
253		// Wrap each MLE-check evaluator so it emits sumcheck-compatible round polynomials.
254		let claims_with_evaluators: [(F, Box<dyn SumcheckRoundEvaluator<F, P>>); 2] = [
255			(eval_claims[0], Box::new(MleToSumCheckEvaluator::new(num_evaluator, eq_tracker))),
256			(eval_claims[1], Box::new(MleToSumCheckEvaluator::new(den_evaluator, eq_tracker))),
257		];
258		let prover = SharedSumcheckProver::new(store, claims_with_evaluators);
259
260		test_frac_add_sumcheck_prove_verify(
261			prover,
262			eval_claims,
263			&eval_point,
264			&num_a,
265			&num_b,
266			&den_a,
267			&den_b,
268		);
269	}
270}