Skip to main content

binius_prover/fold_word/
word_axis.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Contracting the word axis: one field element per bit position.
5
6use std::iter;
7
8use binius_core::word::Word;
9use binius_field::{BinaryField, PackedBinaryField64x1b};
10use binius_math::{
11	FieldBuffer,
12	multilinear::eq::{eq_ind_partial_eval, eq_ind_partial_eval_scalars},
13};
14use binius_utils::rayon::prelude::*;
15
16use super::{CHUNK_SIZE, FoldedWord, LOG_CHUNK_SIZE};
17use crate::bit_matrix::{ColumnSums, RowFoldTables, WEIGHTS_PER_TABLE};
18
19/// Minimum chunks one parallel task folds along the word axis.
20///
21/// A chunk is 64 words and folds in well under a microsecond, which is about what a handoff costs.
22/// Left unbounded, the split reaches one chunk per task and the handoffs dominate.
23///
24/// A floor also caps how far a loop can divide.
25/// A list of `n` chunks splits into at most `n / floor` tasks.
26/// So a floor set too high starves the cores on a short list, and one set too low drowns a long
27/// list in handoffs.
28///
29/// Sixteen chunks is 1024 words, which is the setting that loses at neither end.
30const MIN_CHUNKS_PER_TASK: usize = 16;
31/// A reusable [Method of Four Russians] folder over a fixed evaluation point, contracting the
32/// word axis.
33///
34/// Many word-lists often share one point, so the tables are built once here and reused.
35/// The batched instance fold is that case: every committed word folds against the same point.
36///
37/// The two tables it holds:
38/// * per-byte subset-sum lookups, built from the point's prefix.
39/// * one weight per chunk, built from the point's suffix.
40///
41/// [Method of Four Russians]: <https://en.wikipedia.org/wiki/Method_of_Four_Russians>
42#[derive(Debug)]
43pub struct WordAxisFolder<F: BinaryField> {
44	/// One 256-entry subset-sum table per byte of a word, from the prefix expansion.
45	///
46	/// Table `s` folds the words at positions `s * WEIGHTS_PER_TABLE + t` within a chunk.
47	/// Each such word is weighted by prefix-expansion entry `t` of that group.
48	lookups: RowFoldTables<F, { Word::BYTES }>,
49	/// One weight per chunk of `CHUNK_SIZE` words, from the suffix expansion.
50	suffix_weights: FieldBuffer<F>,
51	/// Base-2 log of the word axis's length, which every folded list fits in.
52	///
53	/// This is the point's own width, stored rather than the length it stands for.
54	/// A point as wide as a `usize` would overflow that length, but never its log.
55	log_n_words: usize,
56}
57
58impl<F: BinaryField> WordAxisFolder<F> {
59	/// Builds the folding tables for `point`.
60	///
61	/// Each later fold takes a list of at most `2^point.len()` words, folded against
62	/// this point.
63	pub fn new(point: &[F]) -> Self {
64		// The point splits into a prefix indexing words within a chunk and a suffix indexing
65		// chunks.
66		let prefix_len = point.len().min(LOG_CHUNK_SIZE);
67		let (prefix, suffix) = point.split_at(prefix_len);
68
69		// One weight per word of a chunk, from the prefix.
70		// A point shorter than one chunk yields fewer weights than a chunk holds, and the table
71		// build reads the rest as zero.
72		// Those zeros pair with the repeated words a short list is filled with, so they add
73		// nothing.
74		let prefix_expansion = eq_ind_partial_eval_scalars(prefix);
75		let lookups = RowFoldTables::new(&prefix_expansion);
76
77		// One weight per chunk of CHUNK_SIZE words, from the suffix.
78		let suffix_weights = eq_ind_partial_eval::<F>(suffix);
79
80		Self {
81			lookups,
82			suffix_weights,
83			log_n_words: point.len(),
84		}
85	}
86
87	/// Folds one word-list against the point.
88	///
89	/// Returns the array whose entry at bit position `b` is
90	///
91	/// ```text
92	/// out[b] = sum_i eq(point, i) * bit_b(words[i])
93	/// ```
94	///
95	/// with a clear bit read as zero and a set bit read as one.
96	///
97	/// This runs sequentially over the list's chunks, so it leaves every other core free.
98	/// It is the right driver for a caller already parallel over many lists.
99	/// A caller with few lists wants the parallel driver below instead.
100	///
101	/// A list shorter than the word axis reads its missing high rows as zero: an absent row's
102	/// weight multiplies nothing, so it contributes nothing to any bit position. Chunks lying
103	/// entirely past the list's end are therefore never visited at all.
104	///
105	/// ## Preconditions
106	///
107	/// * `words.len() <= 1 << point.len()`
108	pub fn fold(&self, words: &[Word]) -> FoldedWord<F> {
109		assert!(words.len() <= 1 << self.log_n_words, "words.len() must not exceed 2^point.len()");
110
111		let (chunks, tail) = words.as_chunks::<CHUNK_SIZE>();
112		let mut folded = [F::ZERO; Word::BITS];
113
114		// Accumulate each chunk's contribution, scaled by that chunk's suffix weight. Weights past
115		// the list's end pair with absent rows, so the zip drops them.
116		for (chunk, &suffix_weight) in iter::zip(chunks, self.suffix_weights.as_ref()) {
117			self.accumulate_chunk(chunk, suffix_weight, &mut folded);
118		}
119
120		self.accumulate_tail(tail, chunks.len(), &mut folded);
121		folded
122	}
123
124	/// Folds one word-list against the point, parallel over that list's chunks.
125	///
126	/// Returns the same array the sequential fold returns, under the same contract.
127	/// The two differ only in how the chunk axis is divided across workers.
128	///
129	/// Reach for this when few lists share the point, so the chunk axis is the only one wide enough
130	/// to divide. A caller folding many lists against one point should instead parallelize across
131	/// the lists and fold each one sequentially.
132	///
133	/// ## Preconditions
134	///
135	/// * `words.len() <= 1 << point.len()`
136	pub fn fold_par(&self, words: &[Word]) -> FoldedWord<F> {
137		assert!(words.len() <= 1 << self.log_n_words, "words.len() must not exceed 2^point.len()");
138
139		let (chunks, tail) = words.as_chunks::<CHUNK_SIZE>();
140
141		// Each chunk contributes to every bit position, scaled by that chunk's suffix weight.
142		// Summing the per-chunk accumulators contracts the word axis.
143		// Weights past the list's end pair with absent rows, so the zip drops them.
144		//
145		// One accumulator per worker, not one per chunk:
146		//
147		//     per chunk : 512 bytes of words in, a 1 KiB accumulator zeroed and merged back out
148		//     per worker: 512 bytes of words in, straight into an accumulator already live
149		//
150		// A merge seeded with a partial that already exists never touches a buffer of zeros.
151		// An identity would zero one accumulator per chunk, then add all 64 elements of it.
152		let mut folded = chunks
153			.par_iter()
154			.zip(self.suffix_weights.as_ref().par_iter())
155			// One item is one chunk of 64 words, so the floor needs no conversion.
156			.with_min_len(MIN_CHUNKS_PER_TASK)
157			.fold(
158				|| [F::ZERO; Word::BITS],
159				|mut acc, (chunk, &suffix_weight)| {
160					self.accumulate_chunk(chunk, suffix_weight, &mut acc);
161					acc
162				},
163			)
164			.reduce_with(|mut lhs, rhs| {
165				for (lhs_i, rhs_i) in iter::zip(&mut lhs, rhs) {
166					*lhs_i += rhs_i;
167				}
168				lhs
169			})
170			// A list with no whole chunks yields no partials at all, and folds to zero.
171			.unwrap_or([F::ZERO; Word::BITS]);
172
173		self.accumulate_tail(tail, chunks.len(), &mut folded);
174		folded
175	}
176
177	/// Folds one chunk of words into the accumulator, scaled by that chunk's weight.
178	///
179	/// Words are 64-bit rows, so a chunk is 64 of them and the columns are the 64 bit positions.
180	/// A word and a 64-bit row of single-bit scalars share one underlier, so the view below is
181	/// free.
182	fn accumulate_chunk(&self, chunk: &[Word; CHUNK_SIZE], weight: F, acc: &mut FoldedWord<F>) {
183		// Reshape the chunk into one contiguous group of eight rows per table.
184		let groups = bytemuck::must_cast_ref::<
185			[Word; CHUNK_SIZE],
186			[[PackedBinaryField64x1b; WEIGHTS_PER_TABLE]; Word::BYTES],
187		>(chunk);
188
189		// Sum every group's contribution before scaling, so the chunk costs one multiply per
190		// column.
191		let mut sums = ColumnSums::zero();
192		self.lookups.fold_into(groups.iter().copied(), &mut sums);
193		sums.add_scaled_to(weight, acc);
194	}
195
196	/// Accumulates the chunk the list ends in, completed with its zero rows.
197	///
198	/// A list whose length is a whole number of chunks has no such chunk, and this does nothing.
199	///
200	/// # Arguments
201	///
202	/// * `tail` - the words after the last whole chunk, fewer than one chunk of them
203	/// * `n_whole_chunks` - how many whole chunks came before, which selects the weight to use
204	/// * `folded` - the accumulator the tail's contribution is added into
205	fn accumulate_tail(&self, tail: &[Word], n_whole_chunks: usize, folded: &mut FoldedWord<F>) {
206		if tail.is_empty() {
207			return;
208		}
209
210		// Rows past the list's end read as zero, which contributes nothing to any bit position.
211		let mut chunk = [Word::ZERO; CHUNK_SIZE];
212		chunk[..tail.len()].copy_from_slice(tail);
213		self.accumulate_chunk(&chunk, self.suffix_weights.get(n_whole_chunks), folded);
214	}
215}
216
217#[cfg(test)]
218mod tests {
219	use binius_math::test_utils::random_scalars;
220	use binius_utils::checked_arithmetics::log2_ceil_usize;
221	use binius_verifier::config::B128;
222	use proptest::prelude::*;
223	use rand::prelude::*;
224
225	use super::*;
226
227	/// Contracts the word axis, leaving one element per bit position.
228	///
229	/// A list shorter than the axis reads its high rows as zero, which weight nothing.
230	fn reference_fold_word_axis<F: BinaryField>(words: &[Word], point: &[F]) -> FoldedWord<F> {
231		assert!(words.len() <= 1 << point.len());
232
233		let eq = eq_ind_partial_eval_scalars(point);
234		let mut out = [F::ZERO; Word::BITS];
235		for (word, &weight) in iter::zip(words, &eq) {
236			for (bit, out_bit) in out.iter_mut().enumerate() {
237				if (word.as_u64() >> bit) & 1 == 1 {
238					*out_bit += weight;
239				}
240			}
241		}
242		out
243	}
244
245	/// Word counts spanning every regime both folds branch on.
246	///
247	/// The folds split their input into whole chunks and a short tail, so the interesting lengths
248	/// sit around those boundaries rather than at round powers of two.
249	fn any_n_words() -> impl Strategy<Value = usize> {
250		0..=4 * CHUNK_SIZE
251	}
252
253	fn words_of(n: usize, seed: u64) -> Vec<Word> {
254		let mut rng = StdRng::seed_from_u64(seed);
255		(0..n).map(|_| Word::from_u64(rng.random())).collect()
256	}
257
258	proptest! {
259		#[test]
260		fn word_axis_fold_matches_the_definition(n_words in any_n_words(), seed: u64) {
261			// The point must cover the list, so its width is the list's rounded-up log.
262			let words = words_of(n_words, seed);
263			let mut rng = StdRng::seed_from_u64(seed ^ 2);
264			let point = random_scalars::<B128>(&mut rng, log2_ceil_usize(words.len()));
265
266			prop_assert_eq!(
267				WordAxisFolder::new(&point).fold_par(&words),
268				reference_fold_word_axis(&words, &point),
269			);
270		}
271
272		#[test]
273		fn word_axis_drivers_agree(n_words in any_n_words(), seed: u64) {
274			// Dividing the chunk axis across workers changes the grouping of the sums, not their
275			// value, because field addition is associative and exact.
276			let words = words_of(n_words, seed);
277			let mut rng = StdRng::seed_from_u64(seed ^ 3);
278			let point = random_scalars::<B128>(&mut rng, log2_ceil_usize(words.len()));
279
280			let folder = WordAxisFolder::new(&point);
281			prop_assert_eq!(folder.fold_par(&words), folder.fold(&words));
282		}
283
284		#[test]
285		fn word_axis_fold_reads_a_short_list_as_zero_padded(
286			log_rows in LOG_CHUNK_SIZE..LOG_CHUNK_SIZE + 3,
287			seed: u64,
288		) {
289			// A list shorter than the word axis must fold as that list zero-padded up to it,
290			// without ever materializing the padding.
291			let n_words = (seed as usize) % (1 << log_rows);
292			let words = words_of(n_words, seed);
293			let mut rng = StdRng::seed_from_u64(seed ^ 5);
294			let point = random_scalars::<B128>(&mut rng, log_rows);
295
296			let mut padded = words.clone();
297			padded.resize(1 << log_rows, Word::ZERO);
298			let expected = reference_fold_word_axis(&padded, &point);
299
300			prop_assert_eq!(WordAxisFolder::new(&point).fold(&words), expected);
301			prop_assert_eq!(WordAxisFolder::new(&point).fold_par(&words), expected);
302		}
303	}
304}