1use std::ops::Range;
6
7use binius_field::{Field, field::FieldOps};
8
9use super::eq::eq_ind_partial_eval_scalars;
10
11const MAX_CHUNK_WIDTH: usize = 16;
15
16#[derive(Debug, Clone, PartialEq, Eq)]
20pub struct SparseBitVector {
21 log_len: usize,
22 indices: Vec<u64>,
23}
24
25impl SparseBitVector {
26 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 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 pub const fn log_len(&self) -> usize {
54 self.log_len
55 }
56
57 pub fn indices(&self) -> &[u64] {
59 &self.indices
60 }
61}
62
63pub 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
99pub 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
125fn 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
137fn 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
142fn 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
171fn 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 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 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 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}