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}