Skip to main content

binius_prover/
bit_matrix.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Folding a matrix of single-bit rows against one weight per row.
5
6use std::{array, iter};
7
8use binius_field::{
9	BinaryField, Divisible, PackedField, UnderlierView, transpose::transpose_square_blocks_array,
10	util::expand_subset_sums_array,
11};
12use binius_verifier::config::B1;
13
14/// Weights one subset-sum table covers.
15///
16/// Eight is the widest group whose lookup index still fits one byte.
17/// One table load then replaces eight conditional additions.
18pub const WEIGHTS_PER_TABLE: usize = 1 << LOG_WEIGHTS_PER_TABLE;
19
20/// Base-2 log of the weights one subset-sum table covers.
21pub const LOG_WEIGHTS_PER_TABLE: usize = 3;
22
23/// The rows one table covers, one bit per scalar.
24pub type RowGroup<PB> = [PB; WEIGHTS_PER_TABLE];
25
26/// Subset-sum tables for folding a matrix of single-bit rows against one weight per row.
27///
28/// The fold contracts the row axis, leaving one field element per column:
29///
30/// ```text
31///     out[b] = sum_r weight[r] * bit_b(row[r])
32/// ```
33///
34/// Taking eight rows at a time turns that inner sum into a single table lookup.
35/// Table `g` covers rows `8g` through `8g + 7` and holds every subset sum of their weights.
36#[derive(Debug, Clone)]
37pub struct RowFoldTables<F, const N_TABLES: usize> {
38	tables: [[F; 1 << WEIGHTS_PER_TABLE]; N_TABLES],
39}
40
41impl<F: BinaryField, const N_TABLES: usize> RowFoldTables<F, N_TABLES> {
42	/// Builds the tables from one weight per row, from the first row onwards.
43	///
44	/// Weights past the end of the slice read as zero.
45	/// Those weight rows past the end of the matrix, which read as zero as well.
46	/// So one table layout serves a chunk the row list does not fill.
47	pub fn new(weights: &[F]) -> Self {
48		let tables = array::from_fn(|group| {
49			// Weights of the eight rows this table covers, zero where the slice has run out.
50			// A group beyond the end of the slice starts at its end, so it copies nothing.
51			let mut group_weights = [F::ZERO; WEIGHTS_PER_TABLE];
52			let start = (group * WEIGHTS_PER_TABLE).min(weights.len());
53			let available = (weights.len() - start).min(WEIGHTS_PER_TABLE);
54			group_weights[..available].copy_from_slice(&weights[start..start + available]);
55
56			// Enumerate all 256 subset sums, so any byte of set bits indexes its sum in one load.
57			expand_subset_sums_array(group_weights)
58		});
59
60		Self { tables }
61	}
62
63	/// Folds each group of eight rows into the column sums.
64	///
65	/// Groups pair with tables in order, so the `g`th group is weighted by rows `8g` onwards.
66	/// An iterator yielding fewer groups leaves the remaining rows out, which reads them as zero.
67	///
68	/// # Preconditions
69	///
70	/// * A row must be one byte wide per table, so the groups cover every column exactly once.
71	#[inline]
72	pub fn fold_into<PB>(
73		&self,
74		groups: impl IntoIterator<Item = RowGroup<PB>>,
75		sums: &mut ColumnSums<F, N_TABLES>,
76	) where
77		PB: PackedField<Scalar = B1> + UnderlierView,
78		PB::Underlier: Divisible<u8>,
79	{
80		// One byte of a row per table is what makes the accumulator line up with the columns.
81		const {
82			assert!(
83				PB::WIDTH == WEIGHTS_PER_TABLE * N_TABLES,
84				"the row width must be one byte per table"
85			);
86		}
87
88		// Pairing here rather than at the call site is what keeps a group with its own weights.
89		for (group, table) in iter::zip(groups, &self.tables) {
90			fold_group(group, table, &mut sums.groups);
91		}
92	}
93}
94
95/// One field element per column of the matrix, summed across row groups.
96///
97/// Reading the groups end to end walks the columns in order.
98#[derive(Debug, Clone, PartialEq, Eq)]
99pub struct ColumnSums<F, const N_TABLES: usize> {
100	// Invariant: entry `j` of group `i` is column `8i + j`, which is flat index `8i + j`.
101	//
102	// So the nesting is the identity permutation, and exists only because an array length of
103	// `WEIGHTS_PER_TABLE * N_TABLES` is not expressible in a const generic on stable.
104	groups: [[F; WEIGHTS_PER_TABLE]; N_TABLES],
105}
106
107impl<F: BinaryField, const N_TABLES: usize> ColumnSums<F, N_TABLES> {
108	/// The sums before any group has been folded in.
109	pub const fn zero() -> Self {
110		Self {
111			groups: [[F::ZERO; WEIGHTS_PER_TABLE]; N_TABLES],
112		}
113	}
114
115	/// The sums, in column order.
116	#[inline]
117	pub const fn as_slice(&self) -> &[F] {
118		self.groups.as_flattened()
119	}
120
121	/// Adds every column's sum into `out`, scaled by one weight.
122	///
123	/// A caller folding a matrix in chunks scales each chunk by its own weight on the way out.
124	///
125	/// # Preconditions
126	///
127	/// * `out` must hold one element per column.
128	#[inline]
129	pub fn add_scaled_to(&self, weight: F, out: &mut [F]) {
130		debug_assert_eq!(out.len(), self.as_slice().len()); // precondition
131
132		for (out_i, &sum) in iter::zip(out, self.as_slice()) {
133			*out_i += sum * weight;
134		}
135	}
136}
137
138impl<F: BinaryField, const N_TABLES: usize> Default for ColumnSums<F, N_TABLES> {
139	fn default() -> Self {
140		Self::zero()
141	}
142}
143
144/// Folds one group of eight rows into the column sums.
145///
146/// Rows arrive one bit per scalar, so a row is one packed element and a column is a scalar index.
147/// Transposing the group exchanges its row axis with the low three bits of the column index:
148///
149/// ```text
150///     before:  element r, bit 8i + j  =  row r, column 8i + j
151///     after:   element j, bit 8i + t  =  row t, column 8i + j
152/// ```
153///
154/// So byte `i` of element `j` carries the eight rows' bits at column `8i + j`.
155/// One lookup of that byte yields those rows' whole contribution to that column.
156#[inline]
157fn fold_group<F, PB, const N_TABLES: usize>(
158	mut group: RowGroup<PB>,
159	table: &[F; 1 << WEIGHTS_PER_TABLE],
160	sums: &mut [[F; WEIGHTS_PER_TABLE]; N_TABLES],
161) where
162	F: BinaryField,
163	PB: PackedField<Scalar = B1> + UnderlierView,
164	PB::Underlier: Divisible<u8>,
165{
166	// The transpose rewrites the group in place, and the caller handed over its copy.
167	transpose_square_blocks_array::<PB, LOG_WEIGHTS_PER_TABLE, WEIGHTS_PER_TABLE>(&mut group);
168
169	for (j, row) in group.iter().enumerate() {
170		// Byte `i` holds this group's bits at column `8i + j`, so it indexes that column's sum.
171		for (i, byte) in Divisible::<u8>::value_iter(row.to_underlier()).enumerate() {
172			sums[i][j] += table[byte as usize];
173		}
174	}
175}
176
177#[cfg(test)]
178mod tests {
179	use binius_field::{Field, PackedBinaryField64x1b, PackedBinaryField128x1b, Random};
180	use binius_math::test_utils::random_scalars;
181	use binius_verifier::config::B128;
182	use rand::prelude::*;
183
184	use super::*;
185
186	// The fold, written straight from its definition: every set bit adds its row's weight to that
187	// bit's column.
188	//
189	//     out[b] = sum_r weight[r] * bit_b(row[r])
190	fn naive_fold<F: BinaryField, PB: PackedField<Scalar = B1>>(
191		rows: &[PB],
192		weights: &[F],
193	) -> Vec<F> {
194		let mut out = vec![F::ZERO; PB::WIDTH];
195		for (row, &weight) in iter::zip(rows, weights) {
196			for (column, bit) in row.iter().enumerate() {
197				if bit == B1::ONE {
198					out[column] += weight;
199				}
200			}
201		}
202		out
203	}
204
205	// One row per lane of the packed type, grouped eight at a time.
206	fn random_groups<PB: PackedField + Random, const N_TABLES: usize>(
207		rng: &mut StdRng,
208	) -> Vec<RowGroup<PB>> {
209		(0..N_TABLES)
210			.map(|_| array::from_fn(|_| PB::random(&mut *rng)))
211			.collect()
212	}
213
214	fn check_matches_naive<PB, const N_TABLES: usize>(seed: u64, n_weights: usize)
215	where
216		PB: PackedField<Scalar = B1> + UnderlierView + Random,
217		PB::Underlier: Divisible<u8>,
218	{
219		let mut rng = StdRng::seed_from_u64(seed);
220
221		// One row per weight the fold covers, and one weight per row the tables cover.
222		let groups = random_groups::<PB, N_TABLES>(&mut rng);
223		let weights = random_scalars::<B128>(&mut rng, n_weights);
224
225		let tables = RowFoldTables::<B128, N_TABLES>::new(&weights);
226		let mut sums = ColumnSums::zero();
227		tables.fold_into(groups.iter().copied(), &mut sums);
228
229		// The naive side needs the rows laid out flat, in the order the tables weight them.
230		let rows = groups.concat();
231		let mut padded = weights;
232		padded.resize(rows.len(), B128::ZERO);
233
234		assert_eq!(
235			sums.as_slice(),
236			naive_fold(&rows, &padded),
237			"fold differs at seed {seed}, {n_weights} weights"
238		);
239	}
240
241	#[test]
242	fn fold_matches_the_definition() {
243		// The two row widths the callers run at: 64-bit words and 128-bit field elements.
244		//
245		//     64 columns  -> 8 tables of 8 rows
246		//     128 columns -> 16 tables of 8 rows
247		check_matches_naive::<PackedBinaryField64x1b, 8>(0, 64);
248		check_matches_naive::<PackedBinaryField128x1b, 16>(1, 128);
249	}
250
251	#[test]
252	fn weights_past_the_end_read_as_zero() {
253		// A row list that does not fill the tables must fold as the same list zero-padded up to
254		// them. This is what lets one table layout serve a partial chunk.
255		//
256		//     weights: [w_0 .. w_20]           rows 21..63 weigh nothing
257		//     padded : [w_0 .. w_20, 0 .. 0]
258		check_matches_naive::<PackedBinaryField64x1b, 8>(2, 21);
259		check_matches_naive::<PackedBinaryField64x1b, 8>(3, 0);
260		check_matches_naive::<PackedBinaryField128x1b, 16>(4, 100);
261	}
262
263	#[test]
264	fn sums_read_out_in_column_order() {
265		let mut rng = StdRng::seed_from_u64(5);
266
267		// Weight row 0 alone, so every column's sum is that weight exactly where row 0 has a set
268		// bit, and zero elsewhere. That pins the read-out order against the row's own bits.
269		let weights = random_scalars::<B128>(&mut rng, 1);
270		let row = PackedBinaryField64x1b::random(&mut rng);
271		let mut groups = [[PackedBinaryField64x1b::default(); WEIGHTS_PER_TABLE]; 8];
272		groups[0][0] = row;
273
274		let tables = RowFoldTables::<B128, 8>::new(&weights);
275		let mut sums = ColumnSums::zero();
276		tables.fold_into(groups.iter().copied(), &mut sums);
277
278		for (column, bit) in row.iter().enumerate() {
279			let expected = if bit == B1::ONE {
280				weights[0]
281			} else {
282				B128::ZERO
283			};
284			assert_eq!(sums.as_slice()[column], expected, "column {column}");
285		}
286	}
287
288	#[test]
289	fn scaling_out_multiplies_every_column() {
290		let mut rng = StdRng::seed_from_u64(6);
291
292		// Folding a chunk and scaling it must equal scaling each column sum by hand, which is what
293		// lets a caller pay one multiply per column instead of one per row.
294		let groups = random_groups::<PackedBinaryField64x1b, 8>(&mut rng);
295		let weights = random_scalars::<B128>(&mut rng, 64);
296		let scale = random_scalars::<B128>(&mut rng, 1)[0];
297
298		let tables = RowFoldTables::<B128, 8>::new(&weights);
299		let mut sums = ColumnSums::zero();
300		tables.fold_into(groups.iter().copied(), &mut sums);
301
302		// Start from a non-zero accumulator, so the addition is exercised and not just the scale.
303		let mut out = random_scalars::<B128>(&mut rng, 64);
304		let before = out.clone();
305		sums.add_scaled_to(scale, &mut out);
306
307		for (column, ((&got, &was), &sum)) in
308			iter::zip(iter::zip(&out, &before), sums.as_slice()).enumerate()
309		{
310			assert_eq!(got, was + sum * scale, "column {column}");
311		}
312	}
313
314	#[test]
315	fn folding_no_groups_leaves_the_sums_at_zero() {
316		// An empty matrix contributes nothing, so the sums stay where they started.
317		let mut rng = StdRng::seed_from_u64(7);
318		let weights = random_scalars::<B128>(&mut rng, 64);
319
320		let tables = RowFoldTables::<B128, 8>::new(&weights);
321		let mut sums = ColumnSums::zero();
322		tables.fold_into(iter::empty::<RowGroup<PackedBinaryField64x1b>>(), &mut sums);
323
324		assert_eq!(sums, ColumnSums::zero());
325	}
326}