1use 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
14pub const WEIGHTS_PER_TABLE: usize = 1 << LOG_WEIGHTS_PER_TABLE;
19
20pub const LOG_WEIGHTS_PER_TABLE: usize = 3;
22
23pub type RowGroup<PB> = [PB; WEIGHTS_PER_TABLE];
25
26#[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 pub fn new(weights: &[F]) -> Self {
48 let tables = array::from_fn(|group| {
49 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 expand_subset_sums_array(group_weights)
58 });
59
60 Self { tables }
61 }
62
63 #[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 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 for (group, table) in iter::zip(groups, &self.tables) {
90 fold_group(group, table, &mut sums.groups);
91 }
92 }
93}
94
95#[derive(Debug, Clone, PartialEq, Eq)]
99pub struct ColumnSums<F, const N_TABLES: usize> {
100 groups: [[F; WEIGHTS_PER_TABLE]; N_TABLES],
105}
106
107impl<F: BinaryField, const N_TABLES: usize> ColumnSums<F, N_TABLES> {
108 pub const fn zero() -> Self {
110 Self {
111 groups: [[F::ZERO; WEIGHTS_PER_TABLE]; N_TABLES],
112 }
113 }
114
115 #[inline]
117 pub const fn as_slice(&self) -> &[F] {
118 self.groups.as_flattened()
119 }
120
121 #[inline]
129 pub fn add_scaled_to(&self, weight: F, out: &mut [F]) {
130 debug_assert_eq!(out.len(), self.as_slice().len()); 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#[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 transpose_square_blocks_array::<PB, LOG_WEIGHTS_PER_TABLE, WEIGHTS_PER_TABLE>(&mut group);
168
169 for (j, row) in group.iter().enumerate() {
170 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 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 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 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 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 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 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 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 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 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 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}