1use std::{array, hint::assert_unchecked, iter};
7
8use binius_compute::Allocator;
9use binius_core::word::Word;
10use binius_field::{BinaryField, PackedField};
11use binius_math::FieldBuffer;
12use binius_utils::rayon::prelude::*;
13
14use super::{lookup::BitWeightTables, output::PackedOutput};
15
16const MIN_WORDS_PER_TASK: usize = 1 << 12;
29
30#[derive(Debug)]
38pub struct BitAxisFolder<F: BinaryField> {
39 tables: BitWeightTables<F>,
41}
42
43impl<F: BinaryField> BitAxisFolder<F> {
44 pub fn new(vec: &[F]) -> Self {
49 Self {
50 tables: BitWeightTables::new(vec),
51 }
52 }
53
54 pub fn fold<P, A>(&self, alloc: &A, words: &[Word]) -> FieldBuffer<P, A::Vec<P>>
58 where
59 P: PackedField<Scalar = F>,
60 A: Allocator,
61 {
62 let mut out = PackedOutput::for_words(alloc, words.len());
63
64 let n_chunks = words.len() / P::WIDTH;
68 let (words_aligned, words_remaining) = words.split_at(n_chunks * P::WIDTH);
69
70 let slots = out.chunk_slots(n_chunks);
71 let word_chunks = words_aligned.par_chunks_exact(P::WIDTH);
72 assert_eq!(slots.len(), word_chunks.len());
73
74 (slots, word_chunks)
75 .into_par_iter()
76 .with_min_len(MIN_WORDS_PER_TASK.div_ceil(P::WIDTH))
78 .for_each(|(slot, word_chunk)| {
79 unsafe { assert_unchecked(word_chunk.len() == P::WIDTH) };
83 slot.write(P::from_scalars(word_chunk.iter().map(|&word| self.tables.fold(word))));
84 });
85
86 unsafe { out.commit_chunks(n_chunks) };
88
89 if !words_remaining.is_empty() {
90 out.push(P::from_scalars(words_remaining.iter().map(|&word| self.tables.fold(word))));
91 }
92
93 out.finish()
94 }
95
96 pub fn fold_bitand_operands<P, A>(
129 &self,
130 alloc: &A,
131 a_words: &[Word],
132 b_words: &[Word],
133 ) -> [FieldBuffer<P, A::Vec<P>>; 3]
134 where
135 P: PackedField<Scalar = F>,
136 A: Allocator,
137 {
138 assert_eq!(a_words.len(), b_words.len());
139
140 let [mut a_out, mut b_out, mut c_out] =
146 array::from_fn(|_| PackedOutput::for_words(alloc, a_words.len()));
147
148 let n_chunks = a_words.len() / P::WIDTH;
152 let (a_aligned, a_remaining) = a_words.split_at(n_chunks * P::WIDTH);
153 let (b_aligned, b_remaining) = b_words.split_at(n_chunks * P::WIDTH);
154
155 let a_slots = a_out.chunk_slots(n_chunks);
156 let b_slots = b_out.chunk_slots(n_chunks);
157 let c_slots = c_out.chunk_slots(n_chunks);
158
159 (
162 a_slots,
163 b_slots,
164 c_slots,
165 a_aligned.par_chunks_exact(P::WIDTH),
166 b_aligned.par_chunks_exact(P::WIDTH),
167 )
168 .into_par_iter()
169 .with_min_len(MIN_WORDS_PER_TASK.div_ceil(P::WIDTH))
171 .for_each(|(a_i, b_i, c_i, a_chunk, b_chunk)| {
172 unsafe {
176 assert_unchecked(a_chunk.len() == P::WIDTH);
177 assert_unchecked(b_chunk.len() == P::WIDTH);
178 }
179 a_i.write(P::from_scalars(a_chunk.iter().map(|&word| self.tables.fold(word))));
181 b_i.write(P::from_scalars(b_chunk.iter().map(|&word| self.tables.fold(word))));
182 c_i.write(P::from_scalars(
184 iter::zip(a_chunk, b_chunk).map(|(&a, &b)| self.tables.fold(a & b)),
185 ));
186 });
187
188 unsafe {
191 a_out.commit_chunks(n_chunks);
192 b_out.commit_chunks(n_chunks);
193 c_out.commit_chunks(n_chunks);
194 }
195
196 if !a_remaining.is_empty() {
198 a_out.push(P::from_scalars(a_remaining.iter().map(|&word| self.tables.fold(word))));
199 b_out.push(P::from_scalars(b_remaining.iter().map(|&word| self.tables.fold(word))));
200 c_out.push(P::from_scalars(
201 iter::zip(a_remaining, b_remaining).map(|(&a, &b)| self.tables.fold(a & b)),
202 ));
203 }
204
205 [a_out, b_out, c_out].map(PackedOutput::finish)
207 }
208}
209
210#[cfg(test)]
211mod tests {
212 use binius_compute::GlobalAllocator;
213 use binius_field::{Field, PackedGhash2x128b, arch::OptimalPackedB128};
214 use binius_math::test_utils::random_scalars;
215 use binius_utils::checked_arithmetics::log2_ceil_usize;
216 use binius_verifier::config::B128;
217 use proptest::prelude::*;
218 use rand::prelude::*;
219
220 use super::*;
221 use crate::fold_word::CHUNK_SIZE;
222
223 fn reference_fold_bit_axis<F: Field, P: PackedField<Scalar = F>>(
228 words: &[Word],
229 weights: &[F],
230 ) -> FieldBuffer<P> {
231 assert_eq!(weights.len(), Word::BITS);
232
233 let log_n = log2_ceil_usize(words.len());
234 let scalars = (0..1 << log_n)
235 .map(|i| {
236 words.get(i).map_or(F::ZERO, |word| {
238 (0..Word::BITS)
239 .filter(|bit| (word.as_u64() >> bit) & 1 == 1)
240 .map(|bit| weights[bit])
241 .sum()
242 })
243 })
244 .collect::<Vec<_>>();
245
246 FieldBuffer::from_values(&scalars)
247 }
248
249 fn any_n_words() -> impl Strategy<Value = usize> {
254 0..=4 * CHUNK_SIZE
255 }
256
257 fn words_of(n: usize, seed: u64) -> Vec<Word> {
258 let mut rng = StdRng::seed_from_u64(seed);
259 (0..n).map(|_| Word::from_u64(rng.random())).collect()
260 }
261
262 fn check_bit_axis_fold<P: PackedField<Scalar = B128>>(words: &[Word], weights: &[B128]) {
271 assert_eq!(
272 BitAxisFolder::new(weights).fold::<P, _>(&GlobalAllocator, words),
273 reference_fold_bit_axis(words, weights),
274 "bit-axis fold differs at P::WIDTH = {}, {} words",
275 P::WIDTH,
276 words.len()
277 );
278 }
279
280 fn check_fused_bitand_fold<P: PackedField<Scalar = B128>>(
281 a_words: &[Word],
282 b_words: &[Word],
283 weights: &[B128],
284 ) {
285 let c_words = iter::zip(a_words, b_words)
286 .map(|(&a, &b)| a & b)
287 .collect::<Vec<_>>();
288 let folder = BitAxisFolder::new(weights);
289
290 let [a, b, c] = folder.fold_bitand_operands::<P, _>(&GlobalAllocator, a_words, b_words);
291 let width = P::WIDTH;
292
293 assert_eq!(a, folder.fold(&GlobalAllocator, a_words), "a differs at width {width}");
294 assert_eq!(b, folder.fold(&GlobalAllocator, b_words), "b differs at width {width}");
295 assert_eq!(c, folder.fold(&GlobalAllocator, &c_words), "c differs at width {width}");
296 }
297
298 proptest! {
299 #[test]
300 fn bit_axis_fold_matches_the_definition(n_words in any_n_words(), seed: u64) {
301 let words = words_of(n_words, seed);
305 let mut rng = StdRng::seed_from_u64(seed ^ 1);
306 let weights = random_scalars::<B128>(&mut rng, Word::BITS);
307
308 check_bit_axis_fold::<OptimalPackedB128>(&words, &weights);
309 check_bit_axis_fold::<PackedGhash2x128b>(&words, &weights);
310 }
311
312 #[test]
313 fn fused_bitand_fold_matches_three_separate_folds(n_words in any_n_words(), seed: u64) {
314 let a_words = words_of(n_words, seed);
319 let b_words = words_of(n_words, seed ^ 0xff);
320
321 let mut rng = StdRng::seed_from_u64(seed ^ 4);
322 let weights = random_scalars::<B128>(&mut rng, Word::BITS);
323
324 check_fused_bitand_fold::<OptimalPackedB128>(&a_words, &b_words, &weights);
325 check_fused_bitand_fold::<PackedGhash2x128b>(&a_words, &b_words, &weights);
326 }
327
328 }
329}