Skip to main content

binius_prover/fold_word/
bit_axis.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Contracting the bit axis: one field element per word.
5
6use 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
16/// Minimum words one parallel task folds along the bit axis.
17///
18/// A packed element is the unit of work here, and it holds as few as one word.
19/// Left unbounded, the split reaches one task per word, and the handoff costs more than the fold.
20///
21/// A floor also caps how far a loop can divide.
22/// A list of `n` words splits into at most `n / floor` tasks.
23/// So the floor must stay below the shortest list this fold runs at, divided by the core count.
24/// Otherwise the split stops before the cores are full.
25///
26/// The shared task-size budgets in the utilities crate are calibrated for wider items.
27/// They land high enough here to breach that cap on a mid-sized list.
28const MIN_WORDS_PER_TASK: usize = 1 << 12;
29
30/// A reusable folder over a fixed vector of bit-index scalars.
31///
32/// The one-shot function above rebuilds its lookup tables on every call.
33/// A caller folding several word-lists against the same scalar vector builds them once here, and
34/// reuses them across folds.
35///
36/// The word axis has a folder of its own, built the same way.
37#[derive(Debug)]
38pub struct BitAxisFolder<F: BinaryField> {
39	/// The subset-sum tables built from the bit-index scalars, shared by every fold.
40	tables: BitWeightTables<F>,
41}
42
43impl<F: BinaryField> BitAxisFolder<F> {
44	/// Builds the folding transform for `vec`.
45	///
46	/// ## Preconditions
47	/// * `vec` contains exactly one scalar per bit of a word
48	pub fn new(vec: &[F]) -> Self {
49		Self {
50			tables: BitWeightTables::new(vec),
51		}
52	}
53
54	/// Folds `words`, mapping each word to the inner product of its bits with the scalar vector.
55	///
56	/// The one-shot function at the top of this module states the exact contract.
57	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		// Partition the words into whole packed-width chunks and a short tail.
65		//
66		//     words:  [ chunk 0 | chunk 1 | ... | chunk n-1 | tail (< P::WIDTH) ]
67		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			// One item is one packed element, so the floor converts from words to items.
77			.with_min_len(MIN_WORDS_PER_TASK.div_ceil(P::WIDTH))
78			.for_each(|(slot, word_chunk)| {
79				// Safety:
80				// - words_aligned has length that is a multiple of P::WIDTH
81				// - words_aligned is split into P::WIDTH chunks
82				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		// Safety: the loop above writes every one of the n_chunks slots exactly once.
87		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	/// Folds the two stored BitAnd operand columns and their derived AND column in one pass.
97	///
98	/// # Overview
99	///
100	/// The BitAnd zerocheck folds three columns of the constraint `A & B = C`.
101	/// On a satisfying witness the third column equals the AND of the first two.
102	/// So this fold reads only the two stored columns and derives the third in registers:
103	///
104	/// ```text
105	///     stream A ──┬──> fold ──> folded A
106	///     stream B ──┼──> fold ──> folded B
107	///                └──> A & B ──> fold ──> folded C   (no third input stream)
108	/// ```
109	///
110	/// # Returns
111	///
112	/// Three folded buffers, in order:
113	/// - the first operand column, folded as a single column would be.
114	/// - the second operand column, folded the same way.
115	/// - the word-by-word AND of the two columns, folded the same way.
116	///
117	/// The AND column is derived in registers and never written to memory.
118	///
119	/// # Performance
120	///
121	/// - Two input streams instead of three.
122	/// - Two register ANDs per word pair replace one memory stream.
123	/// - The bytewise lookup tables stay hot across all three outputs.
124	///
125	/// # Preconditions
126	///
127	/// * The two word-lists have equal length.
128	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		// Padding contract, mirrored from the single-column fold:
141		// the high words up to the next power of two read as zero.
142		// `0 & 0 = 0`, so the derived column stays consistent over that padding.
143		//
144		// One output buffer per folded column, filled through spare capacity.
145		let [mut a_out, mut b_out, mut c_out] =
146			array::from_fn(|_| PackedOutput::for_words(alloc, a_words.len()));
147
148		// Phase 1: partition the inputs into full packed-width chunks and a short tail.
149		//
150		//     words:  [ chunk 0 | chunk 1 | ... | chunk n-1 | tail (< P::WIDTH) ]
151		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		// Phase 2: fold the aligned chunks in parallel.
160		// Each task owns one chunk of both inputs and writes one packed element per output.
161		(
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			// One item is one packed element of each output, so the floor converts from words.
170			.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				// Safety:
173				// - both aligned slices have length n_chunks * P::WIDTH
174				// - both are split into P::WIDTH chunks
175				unsafe {
176					assert_unchecked(a_chunk.len() == P::WIDTH);
177					assert_unchecked(b_chunk.len() == P::WIDTH);
178				}
179				// Fold each stored column by bytewise table lookup.
180				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				// Derive the third column in registers, then fold it the same way.
183				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		// Safety: the loop above writes every one of the n_chunks slots of each output exactly
189		// once.
190		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		// Phase 3: fold the short tail into one final packed element per output.
197		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		// Phase 4: each output zero-pads itself up to the power-of-two capacity.
206		[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	/// Contracts the bit axis, leaving one element per word.
224	///
225	/// A list shorter than a power of two reads its high words as zero, so the padded slots
226	/// fold to zero and the buffer is the next power of two long.
227	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				// Absent high words are zero, and a zero word has no set bits to weight.
237				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	/// Word counts spanning every regime both folds branch on.
250	///
251	/// The folds split their input into whole chunks and a short tail, so the interesting lengths
252	/// sit around those boundaries rather than at round powers of two.
253	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	// The bit-axis folds split their input at a multiple of the packing width, so the width decides
263	// which branches run at all:
264	//
265	//     width 1 : every word is its own chunk, so the short tail is never taken
266	//     width 2 : an odd word count leaves a tail, and the buffer above it is zero-padded
267	//
268	// The optimal packing is one scalar wide on some targets, so pinning only that would leave the
269	// tail and padding paths untested there. Every bit-axis property runs at both widths.
270	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			// Every length, not just the powers of two the old fixture used. A non-power-of-two
302			// list exercises the tail element and the zero padding above it, which is the path
303			// the fixture never reached.
304			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			// The AND-reduction fold derives its third column in registers rather than reading it.
315			// That must equal folding a materialized third column.
316			//
317			//     fused(A, B)  ==  [ fold(A), fold(B), fold(A & B) ]
318			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}