Skip to main content

binius_ip_prover/sumcheck/
factored_multilinear.rs

1// Copyright 2026 The Binius Developers
2
3//! A dense multilinear held as a product of factors over disjoint runs of variables.
4
5use binius_compute::BufferData;
6use binius_field::{Field, PackedField};
7use binius_math::{FieldBuffer, multilinear::fold::fold_highest_var_inplace};
8
9/// A multilinear that factorizes across disjoint runs of its variables.
10///
11/// Some weights are a product of small pieces rather than one table.
12/// An equality indicator over several axes factorizes across them, for instance.
13///
14/// Storing such a weight whole costs the product of the pieces' lengths.
15/// Storing the pieces costs their sum.
16///
17/// ```text
18///     whole:   2^(n_1 + n_2 + n_3) entries
19///     factors: 2^n_1 + 2^n_2 + 2^n_3 entries
20/// ```
21///
22/// The saving is what lets a sumcheck range over a space far too large to materialize, so long as
23/// nothing ever asks for the whole table at once.
24///
25/// # Variable order
26///
27/// Factors are held lowest run first.
28///
29/// A factor over `n` variables owns the next `n` bits of the index, above the bits the factors
30/// before it own. So the last factor owns the highest variables, and is the one a fold consumes
31/// first.
32///
33/// ```text
34///     index = [ factor 0 bits | factor 1 bits | ... | last factor bits ]
35///               lowest                                       highest
36/// ```
37///
38/// # Examples
39///
40/// ```ignore
41/// // A weight over an operand axis, a shift axis, and a value axis.
42/// let mut weight = FactoredMultilinear::new(vec![operand, shift, value]);
43/// let at_index = weight.get(index);
44/// weight.fold_highest_var(challenge);
45/// ```
46/// Where the factors' packed words live.
47///
48/// Defaults to the heap, and the point of the parameter is that it need not be.
49///
50/// A caller proving out of an arena hands over arena-backed factors.
51/// This holds them as they are, rather than forcing a copy onto the heap.
52#[derive(Debug, Clone)]
53pub struct FactoredMultilinear<P: PackedField, Data: BufferData<P> = Vec<P>> {
54	/// The factors still holding variables, lowest run first.
55	///
56	/// A factor is dropped once folding has bound every one of its variables.
57	factors: Vec<FieldBuffer<P, Data>>,
58
59	/// The product of every factor already bound away.
60	///
61	/// A factor that runs out of variables holds one value, which multiplies in here.
62	/// Keeping it separate is what lets the factor list shrink instead of carrying empty buffers.
63	bound: P::Scalar,
64}
65
66impl<P: PackedField, Data: BufferData<P>> FactoredMultilinear<P, Data> {
67	/// Builds a multilinear from its factors, lowest variable run first.
68	///
69	/// A factor with no variables is folded straight into the bound product rather than kept,
70	/// since it holds a value and no axis.
71	pub fn new(factors: impl IntoIterator<Item = FieldBuffer<P, Data>>) -> Self {
72		let mut bound = P::Scalar::ONE;
73		let factors = factors
74			.into_iter()
75			.filter(|factor| {
76				if factor.log_len() == 0 {
77					bound *= factor.get(0);
78					false
79				} else {
80					true
81				}
82			})
83			.collect();
84		Self { factors, bound }
85	}
86
87	/// The number of variables still free.
88	pub fn n_vars(&self) -> usize {
89		self.factors.iter().map(FieldBuffer::log_len).sum()
90	}
91
92	/// The value at one vertex of the hypercube over the free variables.
93	///
94	/// Each factor reads the bits it owns, and the results multiply.
95	/// So a lookup costs one multiplication per factor rather than one per variable.
96	///
97	/// # Panics
98	///
99	/// Panics if the index does not fit the free variables.
100	pub fn get(&self, index: usize) -> P::Scalar {
101		let n_vars = self.n_vars();
102		assert!(
103			n_vars >= usize::BITS as usize || index < 1 << n_vars,
104			"precondition: index {index} must fit {n_vars} variables"
105		);
106
107		let mut value = self.bound;
108		let mut rest = index;
109		for factor in &self.factors {
110			// The factor's own bits are the low ones of what is left.
111			let mask = (1 << factor.log_len()) - 1;
112			value *= factor.get(rest & mask);
113			rest >>= factor.log_len();
114		}
115		value
116	}
117
118	/// Fixes the highest free variable to a value.
119	///
120	/// The highest variable belongs to the last factor, so that is the factor this folds.
121	/// A factor whose variables are all bound holds one value, which moves into the bound product.
122	///
123	/// # Panics
124	///
125	/// Panics if no variable is free.
126	pub fn fold_highest_var(&mut self, challenge: P::Scalar) {
127		let factor = self
128			.factors
129			.last_mut()
130			.expect("precondition: at least one variable must be free");
131		fold_highest_var_inplace(factor, challenge);
132
133		// A factor down to a single value is no longer an axis, so it leaves the list.
134		if factor.log_len() == 0 {
135			self.bound *= factor.get(0);
136			self.factors.pop();
137		}
138	}
139}
140
141#[cfg(test)]
142mod tests {
143	use binius_compute::GlobalAllocator;
144	use binius_field::{Ghash128b, PackedGhash1x128b, Random, arch::OptimalPackedB128};
145	use proptest::prelude::*;
146	use rand::{SeedableRng, prelude::StdRng};
147
148	use super::*;
149
150	/// The product of the factors, materialized in full, as the reference to pin against.
151	fn materialize<P: PackedField>(factors: &[FieldBuffer<P>]) -> FieldBuffer<P> {
152		let n_vars: usize = factors.iter().map(FieldBuffer::log_len).sum();
153		let values = (0..1usize << n_vars)
154			.map(|index| {
155				let mut rest = index;
156				let mut value = P::Scalar::ONE;
157				for factor in factors {
158					let mask = (1 << factor.log_len()) - 1;
159					value *= factor.get(rest & mask);
160					rest >>= factor.log_len();
161				}
162				value
163			})
164			.collect::<Vec<_>>();
165		FieldBuffer::from_values(&values)
166	}
167
168	fn random_factors<P: PackedField>(rng: &mut StdRng, shape: &[usize]) -> Vec<FieldBuffer<P>>
169	where
170		P::Scalar: Random,
171	{
172		shape
173			.iter()
174			.map(|&log_len| {
175				let values = (0..1usize << log_len)
176					.map(|_| P::Scalar::random(&mut *rng))
177					.collect::<Vec<_>>();
178				FieldBuffer::from_values(&values)
179			})
180			.collect()
181	}
182
183	#[test]
184	fn a_factor_with_no_variables_folds_into_the_product() {
185		// Invariant: a factor holding one value is a scalar, not an axis.
186		//
187		// Keeping it in the list would leave a buffer a fold could never consume, and the variable
188		// count would then disagree with what folding can bind.
189		type P = PackedGhash1x128b;
190		let mut rng = StdRng::seed_from_u64(1);
191
192		let factors = random_factors::<P>(&mut rng, &[0, 2, 0]);
193		let expected = materialize(&factors);
194
195		let factored = FactoredMultilinear::new(factors);
196
197		// Only the two-variable factor remains an axis.
198		assert_eq!(factored.n_vars(), 2);
199		for index in 0..4 {
200			assert_eq!(factored.get(index), expected.get(index));
201		}
202	}
203
204	#[test]
205	fn folding_consumes_the_factors_from_the_highest_variable_down() {
206		// Invariant: the last factor owns the highest variables.
207		//
208		// Fixture state: three factors over 1, 2 and 1 variables, so four in total.
209		//
210		//     index bits:  [ f0 : 1 | f1 : 2 | f2 : 1 ]
211		//                    low                 high
212		//
213		// Folding once must bind f2 away entirely, leaving three variables in two factors.
214		type P = PackedGhash1x128b;
215		let mut rng = StdRng::seed_from_u64(2);
216
217		let factors = random_factors::<P>(&mut rng, &[1, 2, 1]);
218		let mut factored = FactoredMultilinear::new(factors.clone());
219		let mut reference = materialize(&factors);
220		assert_eq!(factored.n_vars(), 4);
221
222		let challenge = Ghash128b::random(&mut rng);
223		factored.fold_highest_var(challenge);
224		fold_highest_var_inplace(&mut reference, challenge);
225
226		// The one-variable top factor is gone, so its value now rides the bound product.
227		assert_eq!(factored.n_vars(), 3);
228		for index in 0..8 {
229			assert_eq!(factored.get(index), reference.get(index), "index {index}");
230		}
231	}
232
233	#[test]
234	fn factors_may_live_in_an_arena_rather_than_the_heap() {
235		// Invariant: where a factor's words live is the caller's choice, not this type's.
236		//
237		// A prover working out of an arena builds its weights there.
238		// Forcing them onto the heap would mean a copy per factor.
239		//
240		// The arena exists to avoid exactly that.
241		//
242		// Fixture state: the same two factors twice, one pair on the heap and one in an arena.
243		//
244		//     heap-backed  -> value at each vertex
245		//     arena-backed -> the same value at each vertex
246		//
247		// Folding both and comparing at every step is what shows the storage is invisible.
248		type P = PackedGhash1x128b;
249		let mut rng = StdRng::seed_from_u64(23);
250		let alloc = GlobalAllocator;
251
252		let shape = [2usize, 1];
253		let scalars = shape
254			.iter()
255			.map(|&log_len| {
256				(0..1usize << log_len)
257					.map(|_| Ghash128b::random(&mut rng))
258					.collect::<Vec<_>>()
259			})
260			.collect::<Vec<_>>();
261
262		let mut heap = FactoredMultilinear::<P>::new(
263			scalars
264				.iter()
265				.map(|values| FieldBuffer::<P>::from_values(values)),
266		);
267		let mut arena = FactoredMultilinear::new(
268			scalars
269				.iter()
270				.map(|values| FieldBuffer::<P>::from_values_in(&alloc, values)),
271		);
272
273		assert_eq!(arena.n_vars(), heap.n_vars());
274
275		// Bind every variable, comparing the whole table after each one.
276		while heap.n_vars() > 0 {
277			for index in 0..1usize << heap.n_vars() {
278				assert_eq!(arena.get(index), heap.get(index), "index {index}");
279			}
280			let challenge = Ghash128b::random(&mut rng);
281			heap.fold_highest_var(challenge);
282			arena.fold_highest_var(challenge);
283		}
284		assert_eq!(arena.get(0), heap.get(0));
285	}
286
287	proptest! {
288		/// Folding a factored weight tracks folding the product it stands for, at every step.
289		///
290		/// This is the whole contract: a caller may use the factors in place of the table, and the
291		/// two must agree after any sequence of challenges, not only at the start.
292		#[test]
293		fn folding_tracks_the_materialized_product(
294			shape in prop::collection::vec(0usize..=3, 1..=4),
295			seed: u64,
296		) {
297			type P = OptimalPackedB128;
298			let mut rng = StdRng::seed_from_u64(seed);
299
300			let factors = random_factors::<P>(&mut rng, &shape);
301			let mut factored = FactoredMultilinear::new(factors.clone());
302			let mut reference = materialize(&factors);
303
304			// The materialized reference is the product, so the two must start equal.
305			prop_assert_eq!(factored.n_vars(), reference.log_len());
306			for index in 0..reference.len() {
307				prop_assert_eq!(factored.get(index), reference.get(index));
308			}
309
310			// Bind every variable, comparing after each one.
311			while factored.n_vars() > 0 {
312				let challenge = Ghash128b::random(&mut rng);
313				factored.fold_highest_var(challenge);
314				fold_highest_var_inplace(&mut reference, challenge);
315
316				prop_assert_eq!(factored.n_vars(), reference.log_len());
317				for index in 0..reference.len() {
318					prop_assert_eq!(factored.get(index), reference.get(index));
319				}
320			}
321
322			// Fully bound, both are the single value the whole product folded to.
323			prop_assert_eq!(factored.get(0), reference.get(0));
324		}
325	}
326}