Skip to main content

binius_ip_prover/
claim_fold.rs

1// Copyright 2026 The Binius Developers
2
3//! Proving a fold of evaluation claims on one sparse tensor.
4
5use 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/// An evaluation claim whose point is cut into one run per tensor axis.
18///
19/// The verifier reduces claims over a flat point, since the axes mean nothing to it.
20/// A prover needs the cut, because it builds one weight factor per axis.
21///
22/// Axes run lowest first: the first run owns the lowest index bits.
23#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct AxisClaim<F> {
25	/// The point, one coordinate run per axis, lowest axis first.
26	pub point: Vec<Vec<F>>,
27	/// The value the extension is claimed to take there.
28	pub value: F,
29}
30
31impl<F: Clone> AxisClaim<F> {
32	/// The width of each axis, in variables, lowest axis first.
33	pub fn axes(&self) -> Vec<usize> {
34		self.point.iter().map(Vec::len).collect()
35	}
36
37	/// The same claim over a flat point, as the verifier reduces it.
38	pub fn flatten(&self) -> MultilinearEvalClaim<F> {
39		MultilinearEvalClaim {
40			eval: self.value.clone(),
41			point: self.point.concat(),
42		}
43	}
44}
45
46/// Cuts a flat point into one run per axis, lowest axis first.
47///
48/// This is how a reduced claim regains the structure its inputs had, so a caller may reduce again.
49pub 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
60/// One nonzero of a sparse tensor: a flat index over every axis, and the value carried there.
61///
62/// Axes run lowest first, so the first axis owns the lowest index bits.
63/// Entries at a repeated index add, so a caller need not deduplicate them.
64pub type TensorEntry<F> = (usize, F);
65
66/// Proves a fold of evaluation claims on one tensor, returning the claim they folded to.
67///
68/// The claims all name the same tensor at different points.
69/// One sumcheck reduces them to one point.
70///
71/// What comes out has the same shape as what went in, so a caller can fold again.
72///
73/// # Cost
74///
75/// Per round, one pass over the entries per claim.
76///
77/// Nothing is materialized over the product of the axes.
78/// So the axes may span a space far larger than the entry list.
79///
80/// # Panics
81///
82/// Panics if no claim is given, or if the claims disagree about the axis widths.
83pub 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	// One weight per claim, all riding a single copy of the tensor's entries.
102	//
103	// The weight is an equality indicator over every axis at once, which factorizes across them.
104	// Holding it as one factor per axis is what keeps it off the product of their lengths.
105	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	// One prover over all of them, so the entry list is stored once and folded once.
114	// A prover per claim would hold a copy each, and fold every copy every round.
115	let prover = SparseMultiDenseProductSumcheckProver::new(entries.to_vec(), weights, &sums);
116	let output = batch_prove(vec![prover], channel);
117
118	// The evaluations lead with the tensor's, shared by every claim, then one weight's per claim.
119	let tensor_eval = output.multilinear_evals[0][0];
120
121	// The verifier derives its own weight evaluations, so the tensor's is the one thing it needs.
122	channel.send_one(tensor_eval);
123
124	// The rounds bind the highest variable first, so the point reads back in reverse.
125	let mut point = output.challenges;
126	point.reverse();
127
128	AxisClaim {
129		// Cut the point into one run per axis, matching the shape the claims came in with.
130		point: split_axes(&point, &axes),
131		value: tensor_eval,
132	}
133}
134
135/// The tensor's multilinear extension at one point, evaluated directly from its entries.
136///
137/// This is the reference a folded claim is settled against.
138/// It reads the entries and the point, and nothing from the run that raised the claim.
139///
140/// The indicator is materialized in full, so the cost is one element per vertex of the point.
141/// Call it only where that fits.
142pub 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
151/// Builds a claim asserting the tensor's extension takes the value it actually takes.
152///
153/// The value comes from the entries, so the claim is true by construction.
154pub 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	/// Three axes over 2, 1 and 2 variables, so five in total.
176	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	/// Folds the claims and returns what the verifier made of it.
192	///
193	/// The verifier reduces over flat points.
194	/// So the only fold-specific work here is cutting the reduced point back into axis runs.
195	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		// Invariant: a fold preserves truth, and preserves shape.
218		//
219		// Fixture state: three true claims about one tensor, at three independent points.
220		//
221		//     claim_1 at p_1  -.
222		//     claim_2 at p_2   >-- fold -->  one claim at one new point
223		//     claim_3 at p_3  -'
224		//
225		// The claim that comes out must hold against the tensor read directly.
226		// It must also span the same axes, or it could not be folded again at the next level.
227		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		// A vacuous fixture would let a broken fold pass, so the claims must say something.
234		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		// Invariant: this is the accumulation property, and the whole reason a fold is sound.
249		//
250		// A fold does not check its inputs.
251		// It produces a claim that is false whenever any input was.
252		//
253		// So settling the *output* once, at the root, catches an error anywhere below.
254		//
255		// Fixture state: three true claims, then each one corrupted in turn.
256		//
257		//     before:  claim_i holds
258		//     after:   claim_i's value moved by one, everything else identical
259		//
260		// Either the fold's own assertion rejects, or it returns a claim that does not hold.
261		// Both are correct.
262		//
263		// What must never happen is a true output from a false input.
264		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				// The assertion caught it inside the fold.
276				Err(_) => {}
277				// Or it came through, and the claim it produced must be false.
278				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		// Invariant: one claim is a valid fold, and the reduction still moves the point.
290		//
291		// A tree's leaf level may hand over a single claim.
292		// So the degenerate arity has to work, rather than being a case a caller must avoid.
293		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}