Skip to main content

binius_math/multilinear/
sparse.rs

1// Copyright 2026 The Binius Developers
2
3//! Bit vectors over the hypercube stored as their set bits, and their multilinear extensions.
4
5use std::ops::Range;
6
7use binius_field::{Field, field::FieldOps};
8
9use super::eq::eq_ind_partial_eval_scalars;
10
11/// The widest chunk of coordinates expanded into one equality-indicator tensor.
12///
13/// It bounds every tensor at `2^MAX_CHUNK_WIDTH` elements, however many bits are set.
14const MAX_CHUNK_WIDTH: usize = 16;
15
16/// A bit vector of length `2^log_len`, stored as the indices of its set bits.
17///
18/// Invariant: `indices` is sorted and holds no repeats.
19#[derive(Debug, Clone, PartialEq, Eq)]
20pub struct SparseBitVector {
21	log_len: usize,
22	indices: Vec<u64>,
23}
24
25impl SparseBitVector {
26	/// Sorts the indices and cancels repeated pairs, since a bit set twice is clear.
27	///
28	/// # Preconditions
29	///
30	/// * `log_len` must be less than 64
31	/// * every index must be less than `2^log_len`
32	pub fn new(log_len: usize, mut indices: Vec<u64>) -> Self {
33		assert!(log_len < u64::BITS as usize, "log_len {log_len} must be less than 64");
34		assert!(
35			indices.iter().all(|&index| index >> log_len == 0),
36			"every index must be less than 2^{log_len}"
37		);
38
39		indices.sort_unstable();
40		// Sorting makes equal indices adjacent, so a stack cancels them in pairs.
41		let indices = indices.into_iter().fold(Vec::new(), |mut kept, index| {
42			if kept.last() == Some(&index) {
43				kept.pop();
44			} else {
45				kept.push(index);
46			}
47			kept
48		});
49		Self { log_len, indices }
50	}
51
52	/// Returns the base-2 logarithm of the vector's length.
53	pub const fn log_len(&self) -> usize {
54		self.log_len
55	}
56
57	/// Returns the sorted indices of the set bits.
58	pub fn indices(&self) -> &[u64] {
59		&self.indices
60	}
61}
62
63/// Evaluates the multilinear extension of a bit vector at a point of `bits.log_len()` coordinates.
64///
65/// ```text
66/// Σ_{k in set bits} eq(point, k)
67/// ```
68///
69/// The coordinates are split into contiguous chunks, each expanded once into an
70/// equality-indicator tensor, and each set bit contributes the product of one lookup per chunk.
71/// See [`evaluate_sparse_b1_multilinear_native`] for the faster path over a base field.
72///
73/// ```
74/// # use binius_field::{Field, Ghash128b as B128};
75/// # use binius_math::multilinear::sparse::{SparseBitVector, evaluate_sparse_b1_multilinear};
76/// // Bits 1 and 2 of a length-4 vector, read back at the vertex of index 2.
77/// let bits = SparseBitVector::new(2, vec![2, 1]);
78/// let point = [B128::ZERO, B128::ONE];
79/// assert_eq!(evaluate_sparse_b1_multilinear(&bits, &point), B128::ONE);
80/// ```
81///
82/// # Preconditions
83///
84/// * `point.len()` must equal `bits.log_len()`
85pub fn evaluate_sparse_b1_multilinear<E: FieldOps>(bits: &SparseBitVector, point: &[E]) -> E {
86	let chunks = expand_chunks(bits, point);
87	bits.indices
88		.iter()
89		.map(|&index| {
90			chunks
91				.iter()
92				.map(|(range, tensor)| lookup(index, range, tensor))
93				.reduce(|acc, value| acc * value)
94				.expect("there is at least one chunk")
95		})
96		.sum()
97}
98
99/// Evaluates the multilinear extension of a bit vector natively in the field `F`.
100///
101/// Produces the identical result to [`evaluate_sparse_b1_multilinear`], but multiplies each set
102/// bit's last lookup unreduced with [`WideMul`](binius_field::arithmetic_traits::WideMul) and
103/// reduces the sum once, which the generic path cannot since `E: FieldOps` does not imply it.
104///
105/// # Preconditions
106///
107/// * `point.len()` must equal `bits.log_len()`
108pub fn evaluate_sparse_b1_multilinear_native<F: Field>(bits: &SparseBitVector, point: &[F]) -> F {
109	let chunks = expand_chunks(bits, point);
110	let (last, init) = chunks.split_last().expect("there is at least one chunk");
111	let wide = bits
112		.indices
113		.iter()
114		.map(|&index| {
115			let prefix = init
116				.iter()
117				.map(|(range, tensor)| lookup(index, range, tensor))
118				.product();
119			F::wide_mul(prefix, lookup(index, &last.0, &last.1))
120		})
121		.sum();
122	F::reduce(wide)
123}
124
125/// Splits the point into chunks and expands each into its equality-indicator tensor.
126fn expand_chunks<E: FieldOps>(bits: &SparseBitVector, point: &[E]) -> Vec<(Range<usize>, Vec<E>)> {
127	assert_eq!(point.len(), bits.log_len, "point must have log_len coordinates");
128	chunk_ranges(bits.indices.len(), bits.log_len)
129		.into_iter()
130		.map(|range| {
131			let tensor = eq_ind_partial_eval_scalars(&point[range.clone()]);
132			(range, tensor)
133		})
134		.collect()
135}
136
137/// Reads the tensor entry a chunk's coordinates select from an index.
138fn lookup<E: FieldOps>(index: u64, range: &Range<usize>, tensor: &[E]) -> E {
139	tensor[(index >> range.start) as usize & ((1 << range.len()) - 1)].clone()
140}
141
142/// Splits `log_len` coordinates into contiguous chunks for `n_set` set bits.
143///
144/// The chunk count `k` approximately minimizes the field multiplications
145///
146/// ```text
147/// Σ_i 2^{w_i}  +  n_set · (k − 1)
148/// ```
149///
150/// over widths as equal as they can be, none wider than [`MAX_CHUNK_WIDTH`].
151/// There is always at least one chunk, of width zero when `log_len` is zero.
152fn chunk_ranges(n_set: usize, log_len: usize) -> Vec<Range<usize>> {
153	let min_count = log_len.div_ceil(MAX_CHUNK_WIDTH).max(1);
154	let count = (min_count..=log_len.max(min_count))
155		.min_by_key(|&count| {
156			equal_widths(log_len, count)
157				.map(|width| 1usize << width)
158				.sum::<usize>()
159				+ n_set * (count - 1)
160		})
161		.expect("the range of counts is non-empty");
162	equal_widths(log_len, count)
163		.scan(0, |start, width| {
164			let range = *start..*start + width;
165			*start += width;
166			Some(range)
167		})
168		.collect()
169}
170
171/// Splits `log_len` into `count` widths that differ by at most one.
172fn equal_widths(log_len: usize, count: usize) -> impl Iterator<Item = usize> {
173	(0..count).map(move |i| log_len / count + usize::from(i < log_len % count))
174}
175
176#[cfg(test)]
177mod tests {
178	use std::iter;
179
180	use rand::prelude::*;
181	use rstest::rstest;
182
183	use super::*;
184	use crate::{
185		multilinear::eq::eq_ind,
186		test_utils::{B128, index_to_hypercube_point, random_scalars},
187	};
188
189	fn random_bits(rng: &mut StdRng, n_set: usize, log_len: usize) -> SparseBitVector {
190		let indices = iter::repeat_with(|| rng.random_range(0..1u64 << log_len))
191			.take(n_set)
192			.collect();
193		SparseBitVector::new(log_len, indices)
194	}
195
196	#[rstest]
197	#[case::empty(0, 8, 4)]
198	#[case::zero_vars(1, 0, 1)]
199	#[case::one_bit(1, 16, 8)]
200	#[case::few_bits(8, 12, 4)]
201	#[case::some_bits(64, 12, 3)]
202	#[case::many_bits(1024, 12, 2)]
203	#[case::dense(4096, 12, 1)]
204	fn matches_dense_inner_product(
205		#[case] n_set: usize,
206		#[case] log_len: usize,
207		#[case] n_chunks: usize,
208	) {
209		let mut rng = StdRng::seed_from_u64(0);
210		let bits = random_bits(&mut rng, n_set, log_len);
211		let point = random_scalars::<B128>(&mut rng, log_len);
212
213		// Fewer set bits favor more, narrower chunks; the cases span one to eight.
214		assert_eq!(chunk_ranges(n_set, log_len).len(), n_chunks);
215
216		let tensor = eq_ind_partial_eval_scalars(&point);
217		let expected = bits
218			.indices()
219			.iter()
220			.map(|&index| tensor[index as usize])
221			.sum::<B128>();
222		assert_eq!(evaluate_sparse_b1_multilinear(&bits, &point), expected);
223		assert_eq!(evaluate_sparse_b1_multilinear_native(&bits, &point), expected);
224	}
225
226	#[test]
227	fn wide_point_matches_per_bit_indicator() {
228		let mut rng = StdRng::seed_from_u64(0);
229		let log_len = 40;
230		let bits = random_bits(&mut rng, 3, log_len);
231		let point = random_scalars::<B128>(&mut rng, log_len);
232
233		// Too wide for a dense tensor, so every chunk is capped.
234		assert!(
235			chunk_ranges(3, log_len)
236				.iter()
237				.all(|range| range.len() <= MAX_CHUNK_WIDTH)
238		);
239
240		let expected = bits
241			.indices()
242			.iter()
243			.map(|&index| {
244				// The index is wider than `usize` on 32-bit targets, so the vertex is built from
245				// `u64`.
246				let vertex = (0..log_len)
247					.map(|i| {
248						if (index >> i) & 1 == 1 {
249							B128::ONE
250						} else {
251							B128::ZERO
252						}
253					})
254					.collect::<Vec<_>>();
255				eq_ind(&point, &vertex)
256			})
257			.sum::<B128>();
258		assert_eq!(evaluate_sparse_b1_multilinear(&bits, &point), expected);
259		assert_eq!(evaluate_sparse_b1_multilinear_native(&bits, &point), expected);
260	}
261
262	#[test]
263	fn boolean_point_reads_the_bit() {
264		let mut rng = StdRng::seed_from_u64(0);
265		let log_len = 6;
266		let bits = random_bits(&mut rng, 20, log_len);
267
268		for position in 0..1u64 << log_len {
269			let vertex = index_to_hypercube_point::<B128>(log_len, position as usize);
270			let bit = if bits.indices().contains(&position) {
271				B128::ONE
272			} else {
273				B128::ZERO
274			};
275			assert_eq!(evaluate_sparse_b1_multilinear(&bits, &vertex), bit);
276			assert_eq!(evaluate_sparse_b1_multilinear_native(&bits, &vertex), bit);
277		}
278	}
279
280	#[test]
281	fn repeated_indices_cancel_in_pairs() {
282		let mut rng = StdRng::seed_from_u64(0);
283		let log_len = 4;
284		let uncancelled = vec![3, 5, 3, 7, 3, 5, 9, 9];
285		let bits = SparseBitVector::new(log_len, uncancelled.clone());
286		assert_eq!(bits.indices(), &[3, 7]);
287
288		let point = random_scalars::<B128>(&mut rng, log_len);
289		let naive = uncancelled
290			.iter()
291			.map(|&index| eq_ind(&point, &index_to_hypercube_point(log_len, index as usize)))
292			.sum::<B128>();
293		assert_eq!(evaluate_sparse_b1_multilinear(&bits, &point), naive);
294	}
295
296	#[test]
297	#[should_panic(expected = "every index must be less than 2^4")]
298	fn new_rejects_out_of_range_index() {
299		SparseBitVector::new(4, vec![16]);
300	}
301}