binius_ip_prover/
claim_fold.rs1use binius_field::{Field, PackedField};
6use binius_ip::MultilinearEvalClaim;
7use binius_math::multilinear::eq::{eq_ind_partial_eval, eq_ind_partial_eval_scalars};
8
9use crate::{
10 channel::IPProverChannel,
11 sumcheck::{
12 batch::batch_prove, factored_multilinear::FactoredMultilinear,
13 sparse_dense_product::SparseMultiDenseProductSumcheckProver,
14 },
15};
16
17#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct AxisClaim<F> {
25 pub point: Vec<Vec<F>>,
27 pub value: F,
29}
30
31impl<F: Clone> AxisClaim<F> {
32 pub fn axes(&self) -> Vec<usize> {
34 self.point.iter().map(Vec::len).collect()
35 }
36
37 pub fn flatten(&self) -> MultilinearEvalClaim<F> {
39 MultilinearEvalClaim {
40 eval: self.value.clone(),
41 point: self.point.concat(),
42 }
43 }
44}
45
46pub fn split_axes<F: Clone>(point: &[F], axes: &[usize]) -> Vec<Vec<F>> {
50 let mut runs = Vec::with_capacity(axes.len());
51 let mut rest = point;
52 for &width in axes {
53 let (run, tail) = rest.split_at(width);
54 runs.push(run.to_vec());
55 rest = tail;
56 }
57 runs
58}
59
60pub type TensorEntry<F> = (usize, F);
65
66pub fn prove<F, P>(
84 entries: &[TensorEntry<F>],
85 claims: &[AxisClaim<F>],
86 channel: &mut impl IPProverChannel<F>,
87) -> AxisClaim<F>
88where
89 F: Field,
90 P: PackedField<Scalar = F>,
91{
92 let first = claims
93 .first()
94 .expect("precondition: a fold needs at least one claim");
95 let axes = first.axes();
96 assert!(
97 claims.iter().all(|claim| claim.axes() == axes),
98 "precondition: every claim must span the same axes"
99 );
100
101 let weights = claims
106 .iter()
107 .map(|claim| {
108 FactoredMultilinear::new(claim.point.iter().map(|run| eq_ind_partial_eval::<P>(run)))
109 })
110 .collect::<Vec<_>>();
111 let sums = claims.iter().map(|claim| claim.value).collect::<Vec<_>>();
112
113 let prover = SparseMultiDenseProductSumcheckProver::new(entries.to_vec(), weights, &sums);
116 let output = batch_prove(vec![prover], channel);
117
118 let tensor_eval = output.multilinear_evals[0][0];
120
121 channel.send_one(tensor_eval);
123
124 let mut point = output.challenges;
126 point.reverse();
127
128 AxisClaim {
129 point: split_axes(&point, &axes),
131 value: tensor_eval,
132 }
133}
134
135pub fn evaluate<F: Field>(entries: &[TensorEntry<F>], point: &[Vec<F>]) -> F {
143 let flat = point.concat();
144 let indicator = eq_ind_partial_eval_scalars(&flat);
145 entries
146 .iter()
147 .map(|&(index, value)| value * indicator[index])
148 .sum()
149}
150
151pub fn claim_at<F: Field>(entries: &[TensorEntry<F>], point: Vec<Vec<F>>) -> AxisClaim<F> {
155 let value = evaluate(entries, &point);
156 AxisClaim { point, value }
157}
158
159#[cfg(test)]
160mod tests {
161 use binius_field::{
162 Random,
163 arch::{OptimalB128, OptimalPackedB128},
164 };
165 use binius_ip::batch_eval;
166 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
167 use rand::{SeedableRng, prelude::*};
168
169 use super::*;
170
171 type F = OptimalB128;
172 type P = OptimalPackedB128;
173 type StdChallenger = HasherChallenger<sha2::Sha256>;
174
175 const AXES: [usize; 3] = [2, 1, 2];
177
178 fn random_entries(rng: &mut StdRng, n_entries: usize) -> Vec<TensorEntry<F>> {
179 let n_vars: usize = AXES.iter().sum();
180 (0..n_entries)
181 .map(|_| (rng.random_range(0..1usize << n_vars), F::random(&mut *rng)))
182 .collect()
183 }
184
185 fn random_point(rng: &mut StdRng) -> Vec<Vec<F>> {
186 AXES.iter()
187 .map(|&width| (0..width).map(|_| F::random(&mut *rng)).collect())
188 .collect()
189 }
190
191 fn fold(
196 entries: &[TensorEntry<F>],
197 claims: &[AxisClaim<F>],
198 ) -> Result<AxisClaim<F>, binius_ip::sumcheck::Error> {
199 let mut transcript = ProverTranscript::new(StdChallenger::default());
200 prove::<F, P>(entries, claims, &mut transcript);
201
202 let mut verifier = transcript.into_verifier();
203 let reduced =
204 batch_eval::verify::<F, _>(claims.iter().map(AxisClaim::flatten), &mut verifier)?;
205 verifier
206 .finalize()
207 .expect("the tape must be fully consumed");
208
209 Ok(AxisClaim {
210 point: split_axes(&reduced.point, &claims[0].axes()),
211 value: reduced.eval,
212 })
213 }
214
215 #[test]
216 fn folding_true_claims_yields_a_true_claim() {
217 let mut rng = StdRng::seed_from_u64(1);
228 let entries = random_entries(&mut rng, 20);
229 let claims = (0..3)
230 .map(|_| claim_at(&entries, random_point(&mut rng)))
231 .collect::<Vec<_>>();
232
233 assert!(claims.iter().any(|claim| claim.value != F::ZERO));
235
236 let folded = fold(&entries, &claims).expect("a fold of true claims must verify");
237
238 assert_eq!(folded.axes(), AXES.to_vec(), "the shape must survive the fold");
239 assert_eq!(
240 folded.value,
241 evaluate(&entries, &folded.point),
242 "the folded claim must hold against the tensor"
243 );
244 }
245
246 #[test]
247 fn one_false_claim_makes_the_folded_claim_false() {
248 let mut rng = StdRng::seed_from_u64(2);
265 let entries = random_entries(&mut rng, 20);
266 let honest = (0..3)
267 .map(|_| claim_at(&entries, random_point(&mut rng)))
268 .collect::<Vec<_>>();
269
270 for index in 0..honest.len() {
271 let mut claims = honest.clone();
272 claims[index].value += F::ONE;
273
274 match fold(&entries, &claims) {
275 Err(_) => {}
277 Ok(folded) => assert_ne!(
279 folded.value,
280 evaluate(&entries, &folded.point),
281 "corrupting claim {index} must not fold to a true claim"
282 ),
283 }
284 }
285 }
286
287 #[test]
288 fn a_single_claim_folds_to_a_claim_about_the_same_tensor() {
289 let mut rng = StdRng::seed_from_u64(3);
294 let entries = random_entries(&mut rng, 12);
295 let claim = claim_at(&entries, random_point(&mut rng));
296
297 let folded = fold(&entries, std::slice::from_ref(&claim)).expect("one claim must fold");
298
299 assert_eq!(folded.value, evaluate(&entries, &folded.point));
300 assert_ne!(folded.point, claim.point, "the fold lands on a fresh point");
301 }
302}