Skip to main content

binius_math/multilinear/
fold.rs

1// Copyright 2026 The Binius Developers
2
3//! Fixing the highest variable of a multilinear to a value.
4//!
5//! Fixing one variable of an `n`-variate multilinear leaves an `(n-1)`-variate one:
6//!
7//! ```text
8//! g(X_0, ..., X_{n-2}) = f(X_0, ..., X_{n-2}, r)
9//! ```
10//!
11//! Coefficients are stored with the highest variable selecting which half of the buffer a
12//! coefficient falls in.
13//! So fixing the highest variable pairs the two halves and leaves the result in the first one.
14
15use std::ops::{Deref, DerefMut};
16
17use binius_compute::{Allocator, BufferData, CollectIntoAllocVec};
18use binius_field::{Field, PackedField};
19use binius_utils::{
20	random_access_sequence::RandomAccessSequence,
21	rayon::{
22		prelude::*,
23		task_size::{IndexedParallelIteratorExt, WorkPerItem},
24	},
25};
26
27use crate::{FieldBuffer, FieldVec, line::extrapolate_line};
28
29/// Fixes the highest variable of a multilinear to a value, in place.
30///
31/// ```text
32/// g(X_0, ..., X_{n-2}) = f(X_0, ..., X_{n-2}, r)
33/// ```
34///
35/// The result occupies the first half of the buffer.
36/// The buffer then reports one variable fewer.
37///
38/// ## Preconditions
39///
40/// * the buffer must have at least one variable
41pub fn fold_highest_var_inplace<P: PackedField, Data: BufferData<P>>(
42	values: &mut FieldBuffer<P, Data>,
43	scalar: P::Scalar,
44) {
45	// Each scalar of the result costs one multiplication.
46	// Broadcasting the challenge once lets every packed word reuse the same multiplier.
47	let broadcast_scalar = P::broadcast(scalar);
48	{
49		// The two halves are the multilinear specialized to 0 and to 1 on the highest variable.
50		let mut split = values.split_half_mut();
51		let (mut lo, mut hi) = split.halves();
52		// Interpolate the line through each pair at the challenge, overwriting the low half.
53		(lo.as_mut(), hi.as_mut())
54			.into_par_iter()
55			.with_min_task(WorkPerItem::FieldMuls)
56			.for_each(|(lo_i, hi_i)| {
57				*lo_i = extrapolate_line(*lo_i, *hi_i, broadcast_scalar);
58			});
59	}
60
61	// The result occupies a prefix, so the truncation drops the scalars past it.
62	values.truncate(values.log_len() - 1);
63}
64
65/// Fixes the highest variable of a multilinear to a value, writing into memory from an allocator.
66///
67/// ```text
68/// g(X_0, ..., X_{n-2}) = f(X_0, ..., X_{n-2}, r)
69/// ```
70///
71/// Each output coefficient interpolates the line through one pair of input coefficients.
72/// The input is left untouched.
73///
74/// Use this when the input is borrowed or must be preserved.
75/// Otherwise prefer the form that overwrites the input.
76///
77/// ## Preconditions
78///
79/// * the buffer must have at least one variable
80pub fn fold_highest_var<A: Allocator, P: PackedField, Data: Deref<Target = [P]>>(
81	alloc: &A,
82	values: &FieldBuffer<P, Data>,
83	scalar: P::Scalar,
84) -> FieldVec<P, A> {
85	assert!(values.log_len() > 0, "precondition: buffer must have at least one variable");
86
87	// The two halves are the multilinear specialized to 0 and to 1 on the highest variable.
88	let broadcast_scalar = P::broadcast(scalar);
89	let (lo, hi) = values.split_half();
90
91	// Interpolate the line through each pair at the challenge directly into a fresh buffer
92	// drawn from the allocator.
93	let data = (lo.as_ref(), hi.as_ref())
94		.into_par_iter()
95		.with_min_task(WorkPerItem::FieldMuls)
96		.map(|(&lo_i, &hi_i)| extrapolate_line(lo_i, hi_i, broadcast_scalar))
97		.collect_into_alloc_vec(alloc);
98	FieldBuffer::new(values.log_len() - 1, data)
99}
100
101/// Overwrites a buffer with the high fold of a bit sequence by a tensor.
102///
103/// The bits are the coefficients of a multilinear whose values are all zero or one.
104/// Each output vertex fixes that polynomial's low-indexed variables to that vertex.
105/// What remains is then paired with the tensor.
106///
107/// This runs on one thread.
108///
109/// ## Preconditions
110///
111/// * the bit count must be a power of two
112/// * the bit count must equal the output length times the tensor length
113pub fn binary_fold_high<P, DataOut, DataIn>(
114	values: &mut FieldBuffer<P, DataOut>,
115	tensor: &FieldBuffer<P, DataIn>,
116	bits: &(impl RandomAccessSequence<bool> + Sync),
117) where
118	P: PackedField,
119	DataOut: DerefMut<Target = [P]>,
120	DataIn: Deref<Target = [P]>,
121{
122	assert!(bits.len().is_power_of_two(), "precondition: bits length must be a power of two");
123
124	let values_log_len = values.log_len();
125	// Below one packed word the buffer still occupies a whole word, so only the live lanes count.
126	let width = P::WIDTH.min(values.len());
127
128	assert_eq!(
129		1 << (values_log_len + tensor.log_len()),
130		bits.len(),
131		"precondition: bits length must equal values length times tensor length"
132	);
133
134	values
135		.iter_packed_mut()
136		.enumerate()
137		.for_each(|(i, packed)| {
138			*packed = P::from_scalars((0..width).map(|j| {
139				// The output vertex this lane holds, as an index into the bits' low variables.
140				let scalar_index = i << P::LOG_WIDTH | j;
141				let mut acc = P::Scalar::ZERO;
142
143				// Sum the tensor entries whose bit is set, over the bits' high variables.
144				// Multiplication by a bit is a selection, so no field multiplication is needed.
145				for (k, tensor_packed) in tensor.iter_packed().enumerate() {
146					for (l, tensor_scalar) in tensor_packed.iter().take(tensor.len()).enumerate() {
147						let tensor_scalar_index = k << P::LOG_WIDTH | l;
148						if bits.get(tensor_scalar_index << values_log_len | scalar_index) {
149							acc += tensor_scalar;
150						}
151					}
152				}
153
154				acc
155			}));
156		});
157}
158
159#[cfg(test)]
160mod tests {
161	use std::iter::repeat_with;
162
163	use binius_compute::GlobalAllocator;
164	use binius_utils::rayon::task_size::min_len_for_work;
165	use proptest::prelude::*;
166	use rand::prelude::*;
167
168	use super::*;
169	use crate::{
170		multilinear::eq::eq_ind_partial_eval,
171		test_utils::{B128, Packed128b, random_field_buffer, random_scalars},
172	};
173
174	type P = Packed128b;
175	type F = B128;
176
177	// The packing width is four scalars, so this range straddles it in both directions.
178	const MAX_VARS: usize = 8;
179
180	#[test]
181	fn fold_splits_above_the_task_threshold() {
182		let mut rng = StdRng::seed_from_u64(0);
183
184		// Invariant: a fold splits only once each half holds two minimum tasks.
185		// Below that it runs inline, leaving the parallel path unexercised.
186		//
187		//     words in one half = 2^(n_vars - 1) / scalars per word
188		//     smallest n_vars with words in one half >= 2 * minimum
189		//
190		// Every other fold test is smaller, so this one covers the split.
191		let min_len = min_len_for_work(WorkPerItem::FieldMuls);
192		let n_vars = (2 * min_len * P::WIDTH).next_power_of_two().ilog2() as usize + 1;
193		let half = 1 << (n_vars - 1);
194		let original = random_field_buffer::<P>(&mut rng, n_vars);
195		let challenge = random_scalars::<F>(&mut rng, 1)[0];
196
197		let mut folded = original.clone();
198		fold_highest_var_inplace(&mut folded, challenge);
199		assert_eq!(folded.log_len(), n_vars - 1);
200
201		// Scalar reference: each output interpolates one (lo, hi) pair at the challenge.
202		for i in 0..half {
203			let expected = extrapolate_line(original.get(i), original.get(i | half), challenge);
204			assert_eq!(folded.get(i), expected, "mismatch at index {i}");
205		}
206	}
207
208	#[test]
209	fn folding_a_slice_backed_buffer_matches_folding_an_owned_one() {
210		let mut rng = StdRng::seed_from_u64(0);
211
212		// Invariant: a fold shrinks its buffer, which each backing store does its own way.
213		// A vector drops its tail, a mutable slice re-slices itself.
214		// Both must leave the same coefficients behind.
215		//
216		// Fixture state: one 5-variable buffer, folded twice by the same challenge.
217		//
218		//     owned store      [ c_0 ... c_31 ]  -> [ c'_0 ... c'_15 ]
219		//     slice store      [ c_0 ... c_31 ]  -> [ c'_0 ... c'_15 ]
220		let original = random_field_buffer::<P>(&mut rng, 5);
221		let scalar = random_scalars::<F>(&mut rng, 1)[0];
222
223		let mut expected = original.clone();
224		fold_highest_var_inplace(&mut expected, scalar);
225
226		let mut owned = original;
227		let mut slice = owned.as_mut_view();
228		fold_highest_var_inplace(&mut slice, scalar);
229
230		assert_eq!(slice.log_len(), 4);
231		assert_eq!(slice, expected.as_mut_view());
232	}
233
234	proptest! {
235		#[test]
236		fn the_two_folds_agree(
237			n_vars in 1..=MAX_VARS,
238			seed: u64,
239		) {
240			let mut rng = StdRng::seed_from_u64(seed);
241			let original = random_field_buffer::<P>(&mut rng, n_vars);
242			let scalar = random_scalars::<F>(&mut rng, 1)[0];
243
244			// Out of place leaves the input alone and returns a fresh half-size buffer.
245			let out_of_place = fold_highest_var(&GlobalAllocator, &original, scalar);
246
247			// In place overwrites the input's first half and shrinks it.
248			let mut in_place = original;
249			fold_highest_var_inplace(&mut in_place, scalar);
250
251			prop_assert_eq!(out_of_place.log_len(), n_vars - 1);
252			prop_assert_eq!(out_of_place, in_place);
253		}
254
255		#[test]
256		fn the_binary_fold_matches_folding_the_widened_bits(
257			dest_vars in 0..=6usize,
258			tensor_vars in 0..=4usize,
259			seed: u64,
260		) {
261			let mut rng = StdRng::seed_from_u64(seed);
262			let point = random_scalars::<F>(&mut rng, tensor_vars);
263			let tensor = eq_ind_partial_eval::<P>(&point);
264
265			// The bit count is the product of the two lengths, as the precondition demands.
266			let bits = repeat_with(|| rng.random())
267				.take(1 << (dest_vars + tensor_vars))
268				.collect::<Vec<bool>>();
269
270			let mut folded = FieldBuffer::<P>::zeros(dest_vars);
271			binary_fold_high(&mut folded, &tensor, &bits.as_slice());
272
273			// Reference: widen the bits to field elements and fold the tensor's variables off.
274			let scalars = bits
275				.iter()
276				.map(|&bit| if bit { F::ONE } else { F::ZERO })
277				.collect::<Vec<F>>();
278			let mut reference = FieldBuffer::<P>::from_values(&scalars);
279			for &coord in point.iter().rev() {
280				fold_highest_var_inplace(&mut reference, coord);
281			}
282
283			prop_assert_eq!(folded, reference);
284		}
285	}
286}