1use 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
13pub type LayerProver<'a, A, F, P> =
18 SharedMleCheckProver<'a, A, F, P, Box<dyn MleCheckRoundEvaluator<F, P> + 'a>>;
19
20pub 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 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
70pub 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 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 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 prover_transcript
150 .message()
151 .write_scalar_slice(&prover_evals);
152
153 let mut verifier_transcript = prover_transcript.into_verifier();
155 let sumcheck_output =
156 batch_verify(n_vars, 3, &eval_claims, &mut verifier_transcript).unwrap();
158
159 let mut reduced_eval_point = sumcheck_output.challenges.clone();
161 reduced_eval_point.reverse();
162
163 let multilinear_evals: Vec<F> = verifier_transcript.message().read_vec(4).unwrap();
165
166 let eq_ind_eval = eq_ind(eval_point, &reduced_eval_point);
168
169 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 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 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 let eq_tracker = store.register_eq_tracker(&eval_point);
253 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}