Skip to main content

binius_prover/protocols/bitand/
sumcheck_round_messages.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{array, borrow::Cow, iter};
5
6use binius_core::word::Word;
7use binius_field::{
8	BinaryField, BinaryField1b as B1, ExtensionField, Field, PackedField, PackedRijndael64x8b,
9	Rijndael8b as B8, WideMul, util::expand_subset_sums_array,
10};
11use binius_math::{BinarySubspace, multilinear::eq::eq_ind_partial_eval};
12use binius_utils::rayon::{self, iter::Either, prelude::*};
13use binius_verifier::{
14	config::PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES, protocols::bitand::ROWS_PER_HYPERCUBE_VERTEX,
15};
16use itertools::izip;
17
18use super::ntt_lookup::NTTLookup;
19
20/// Number of big field zerocheck challenges whose equality indicator is expanded per window.
21///
22/// This value controls a memory-versus-parallelism trade-off.
23///
24/// Doubling it doubles the number of precomputed multiplication tables built for one call.
25/// Each table holds 256 field elements.
26/// That is real, measurable memory cost.
27///
28/// Doubling it also doubles the number of words handled together in one parallel chunk of the
29/// round message computation.
30/// That makes the parallel work coarser-grained.
31///
32/// Halving it shrinks both costs.
33/// The price is building more, smaller tables and chunks.
34const N_FIXED_LARGE_CHALLENGES: usize = 4;
35
36/// Generates a univariate polynomial for the sumcheck protocol in AND constraint reduction.
37///
38/// Let our oblong polynomials be A(Z, X₀, ...), B(Z, X₀, ...), and C(Z, X₀, ...)
39///
40/// Let our zerocheck challenges be (r₀, ...)
41///
42/// It turns out that the first k zerocheck challenges can actually be deterministic, since our
43/// polynomials have 1-bit coefficients as long as their tensor product expansion is an
44/// F2-linearly-independent set.
45///
46/// Note: Deterministic here means that the first k zerocheck challenges are a compile-time
47/// agreed-upon parameter to the proof, and not sampled randomly by the verifier
48///
49///
50/// We choose k=3 because we want them to be in a field isomorphic to the 8-bit NTT domain field
51///
52/// Computes a univariate polynomial:
53/// R₀(Z) = ∑_{X₀,...,Xₙ₋₁ ∈ {0,1}} (A·B - C)·eq(X₀,...,Xₙ₋₁; r₀,...,rₙ₋₁)
54///
55/// This is zero at every point on the hypercube IFF A*B-C evaluates to zero at (r₀,...,rₙ₋₁)
56/// for every Z on the univariate domain. Since R₀(Z) is 0 on the univariate domain, the prover
57/// sends only enough values such that the verifier learns a domain of evaluations of size >
58/// deg(R₀(Z))
59///
60/// The product constraint column C is not an input.
61/// Each C word is derived in registers as the AND of the matching A and B words.
62/// A satisfying witness makes that derivation exact on every row.
63/// So no third column is ever built or streamed.
64///
65/// # Arguments
66///
67/// * `log_words` - Base-2 logarithm of the constraint axis's length
68/// * `a_words` - First multiplicand (a) as a one-bit oblong multilinear polynomial
69/// * `b_words` - Second multiplicand (b) as a one-bit oblong multilinear polynomial
70/// * `eq_ind_big_field_challenges` - Partial equality indicator evaluations for big field variables
71/// * `prover_message_domain` - The NTT domain subspace (dimension `SKIPPED_VARS + 1`) from which
72///   the low-degree-extension lookup table is built internally
73///
74/// # Preconditions
75///
76/// * The two columns have equal length, at most `1 << log_words`. They need not fill the constraint
77///   axis: a shorter column has its remaining rows read as zero. Such a row forces the derived `C =
78///   A & B` to zero as well, so `A * B - C` vanishes on it at every point of the univariate domain
79///   and it adds nothing to the message.
80///
81/// # Returns
82///
83/// The evaluations of R₀(Z), a univariate polynomial of degree at most 2*(|D| - 1) where |D| is the
84/// domain size, on another, disjoint |D|-sized domain. This allows the verifier to construct R₀(Z),
85/// since it must equal zero on D.
86///
87/// # Type Parameters
88///
89/// * `F` - The challenge field type (must be a binary field)
90///
91/// # Panics
92///
93/// Panics if any of the following don't hold:
94/// - `big_field_challenges.len() == log_words.saturating_sub(N_FIXED_SMALL_CHALLENGES)`
95/// - `a_words.len() == b_words.len()`
96/// - `a_words.len() <= 1 << log_words`
97pub fn univariate_round_message_extension_domain<F>(
98	log_words: usize,
99	a_words: &[Word],
100	b_words: &[Word],
101	big_field_challenges: &[F],
102	prover_message_domain: &BinarySubspace<B8>,
103) -> [F; ROWS_PER_HYPERCUBE_VERTEX]
104where
105	F: BinaryField + From<B8>,
106{
107	const N_FIXED_SMALL_CHALLENGES: usize = PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len();
108
109	const LOG_CHUNK_SIZE: usize = N_FIXED_SMALL_CHALLENGES + N_FIXED_LARGE_CHALLENGES;
110
111	const CHUNK_SIZE: usize = 1 << LOG_CHUNK_SIZE;
112
113	assert_eq!(big_field_challenges.len(), log_words.saturating_sub(N_FIXED_SMALL_CHALLENGES));
114	assert_eq!(a_words.len(), b_words.len());
115	assert!(a_words.len() <= 1 << log_words);
116
117	let ntt_lookup = tracing::debug_span!("Compute univariate LDE table")
118		.in_scope(|| NTTLookup::new(prover_message_domain));
119
120	let eq_ind_small: [_; 1 << N_FIXED_SMALL_CHALLENGES] =
121		eq_ind_partial_eval::<B8>(&PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES)
122			.iter_scalars()
123			.map(PackedRijndael64x8b::broadcast)
124			.collect::<Vec<_>>()
125			.try_into()
126			.expect("PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len() == N_FIXED_SMALL_CHALLENGES");
127
128	let (eq_ind_fixed_large, extra_challenges) = eq_ind_fixed_large(big_field_challenges);
129	let outer_weight_mul_maps = eq_ind_fixed_large.map(B8ToExtMulMap::new);
130	let eq_ind_extra = eq_ind_partial_eval::<F>(extra_challenges);
131
132	let a_chunks_iter = padded_chunks::<CHUNK_SIZE>(a_words);
133	let b_chunks_iter = padded_chunks::<CHUNK_SIZE>(b_words);
134
135	(a_chunks_iter, b_chunks_iter)
136		.into_par_iter()
137		.map(|(a_chunk, b_chunk)| {
138			// Reshape the chunk arrays into arrays of arrays
139			let [a_subchunks, b_subchunks] = [&a_chunk, &b_chunk].map(|chunk| {
140				bytemuck::must_cast_ref::<
141					[Word; CHUNK_SIZE],
142					[[Word; 1 << N_FIXED_SMALL_CHALLENGES]; 1 << N_FIXED_LARGE_CHALLENGES],
143				>(chunk)
144			});
145
146			let mut acc = [F::ZERO; ROWS_PER_HYPERCUBE_VERTEX];
147			for (a_subchunk, b_subchunk, outer_weight) in
148				izip!(a_subchunks, b_subchunks, &outer_weight_mul_maps)
149			{
150				let mut summed_ntt = <PackedRijndael64x8b as WideMul>::Output::default();
151				for (&a_i, &b_i, inner_weight) in izip!(a_subchunk, b_subchunk, &eq_ind_small) {
152					let c_i = a_i & b_i;
153
154					// Compute the low-degree extension of each column via the lookup table.
155					let a_lde = ntt_lookup.ntt(a_i);
156					let b_lde = ntt_lookup.ntt(b_i);
157					let c_lde = ntt_lookup.ntt(c_i);
158
159					// Compute the weighted composition of the LDE values.
160					summed_ntt +=
161						PackedRijndael64x8b::wide_mul(a_lde * b_lde - c_lde, *inner_weight);
162				}
163
164				let summed_ntt_reduced = PackedRijndael64x8b::reduce(summed_ntt);
165				for (acc_i, summed_ntt_i) in iter::zip(&mut acc, summed_ntt_reduced.into_iter()) {
166					*acc_i += outer_weight.call(summed_ntt_i);
167				}
168			}
169			acc
170		})
171		.zip(eq_ind_extra.as_ref())
172		.map(|(mut acc, eq_weight)| {
173			for acc_i in &mut acc {
174				*acc_i *= eq_weight;
175			}
176			acc
177		})
178		.reduce(
179			|| [F::ZERO; ROWS_PER_HYPERCUBE_VERTEX],
180			|mut lhs, rhs| {
181				for (lhs_i, rhs_i) in iter::zip(&mut lhs, rhs) {
182					*lhs_i += rhs_i;
183				}
184				lhs
185			},
186		)
187}
188
189/// The equality indicator expansion of the first `N_FIXED_LARGE_CHALLENGES` big field zerocheck
190/// challenges, along with the challenges past them.
191///
192/// The expansion weights the windows' subchunks, while the extra challenges weight the windows
193/// themselves. A challenge vector shorter than `N_FIXED_LARGE_CHALLENGES` is zero-extended to that
194/// length, which leaves no extra challenges: a column that short occupies a single window, whose
195/// unused index bits are the ones the zero challenges cover.
196fn eq_ind_fixed_large<F: Field>(
197	big_field_challenges: &[F],
198) -> ([F; 1 << N_FIXED_LARGE_CHALLENGES], &[F]) {
199	if big_field_challenges.len() < N_FIXED_LARGE_CHALLENGES {
200		let eq_ind_fixed_large = eq_ind_partial_eval::<F>(big_field_challenges);
201		let mut eq_ind_fixed_large_padded = [F::ZERO; 1 << N_FIXED_LARGE_CHALLENGES];
202		eq_ind_fixed_large_padded[..eq_ind_fixed_large.len()]
203			.copy_from_slice(eq_ind_fixed_large.as_ref());
204
205		(eq_ind_fixed_large_padded, &[][..])
206	} else {
207		let (fixed_large_challenges, extra_challenges) =
208			big_field_challenges.split_at(N_FIXED_LARGE_CHALLENGES);
209		let fixed_large_challenges: [_; N_FIXED_LARGE_CHALLENGES] = fixed_large_challenges
210			.try_into()
211			.expect("big_field_challenges.len() >= N_FIXED_LARGE_CHALLENGES");
212
213		let eq_ind_fixed_large: [_; 1 << N_FIXED_LARGE_CHALLENGES] =
214			eq_ind_partial_eval::<F>(&fixed_large_challenges)
215				.as_ref()
216				.try_into()
217				.expect("fixed_large_challenges.len() == N_FIXED_LARGE_CHALLENGES");
218
219		(eq_ind_fixed_large, extra_challenges)
220	}
221}
222
223/// The words as chunks of `CHUNK_SIZE`, the last one padded out if the words run out partway
224/// through it.
225///
226/// A chunk the words fill is borrowed in place; only the partial one is copied. The copy zero-
227/// extends the words to the constraint axis's length, then repeats that axis across the rest of the
228/// chunk. A zero row forces the derived `C = A & B` to zero as well, so `A * B - C` vanishes on it
229/// at every point of the univariate domain and it adds nothing to the round message.
230///
231/// Repetition is a no-op unless the whole axis is shorter than one chunk. In that case the chunk's
232/// index bits past the axis carry the fixed small zerocheck challenges, which are non-zero: summing
233/// a non-zero eq challenge over a duplicated coordinate gives back exactly one copy, whereas
234/// leaving those slots zero would scale the axis by `1 + r`. Repetition is what keeps such a
235/// column's round message equal to the message the verifier reconstructs over the axis's own
236/// variables.
237///
238/// # Preconditions
239///
240/// * `CHUNK_SIZE` is a power of two
241/// * The constraint axis is `words.len()` rounded up to a power of two
242fn padded_chunks<const CHUNK_SIZE: usize>(
243	words: &[Word],
244) -> impl IndexedParallelIterator<Item = Cow<'_, [Word; CHUNK_SIZE]>> {
245	let chunks_iter = words.par_chunks_exact(CHUNK_SIZE);
246	let tail = chunks_iter.remainder();
247
248	let chunks_iter = chunks_iter.map(|chunk| {
249		Cow::Borrowed(
250			<&[Word; CHUNK_SIZE]>::try_from(chunk)
251				.expect("chunks_exact produces slices with len CHUNK_SIZE"),
252		)
253	});
254
255	if tail.is_empty() {
256		Either::Right(chunks_iter)
257	} else {
258		// The axis's rows within one chunk. Both are powers of two, so this divides `CHUNK_SIZE`.
259		let axis_rows = words.len().next_power_of_two().min(CHUNK_SIZE);
260
261		let mut tail_padded = [Word::ZERO; CHUNK_SIZE];
262		tail_padded[..tail.len()].copy_from_slice(tail);
263
264		let (axis, rest) = tail_padded.split_at_mut(axis_rows);
265		for copy in rest.chunks_exact_mut(axis_rows) {
266			copy.copy_from_slice(axis);
267		}
268
269		Either::Left(chunks_iter.chain(rayon::iter::once(Cow::Owned(tail_padded))))
270	}
271}
272
273/// Represents a precomputed multiplication map by an extension field constant for
274/// [`B8`].
275///
276/// Multiplication by a constant for a binary field is an $\mathbb{F}_2$-linear transform. For small
277/// inputs, such as $\mathbb{F}_{2^8}$ elements, this can be represented by a small lookup table.
278struct B8ToExtMulMap<F> {
279	lookup: [F; 256],
280}
281
282impl<F: BinaryField + From<B8>> B8ToExtMulMap<F> {
283	fn new(val: F) -> Self {
284		let basis_images: [F; 8] = array::from_fn(|i| {
285			let basis = <B8 as ExtensionField<B1>>::basis(i);
286			F::from(basis) * val
287		});
288		Self {
289			lookup: expand_subset_sums_array(basis_images),
290		}
291	}
292
293	#[inline]
294	const fn call(&self, input: B8) -> F {
295		self.lookup[input.val() as usize]
296	}
297}
298
299#[cfg(test)]
300mod test {
301	use std::iter::repeat_with;
302
303	use binius_compute::GlobalAllocator;
304	use binius_field::{Field, Ghash128b as B128, Random};
305	use binius_math::{BinarySubspace, FieldBuffer, univariate::EvaluationDomain};
306	use binius_utils::checked_arithmetics::log2_ceil_usize;
307	use binius_verifier::protocols::bitand::SKIPPED_VARS;
308	use rand::prelude::*;
309
310	use super::*;
311	use crate::fold_word::BitAxisFolder;
312
313	fn random_words(log_num_words: usize, mut rng: impl Rng) -> Vec<Word> {
314		repeat_with(|| Word(rng.random()))
315			.take(1 << log_num_words)
316			.collect()
317	}
318
319	// Sends the sum claim from first multilinear round (second overall round)
320	pub fn sum_claim<BF: Field + From<B128>>(
321		first_col: &FieldBuffer<BF>,
322		second_col: &FieldBuffer<BF>,
323		third_col: &FieldBuffer<BF>,
324		eq_ind: &FieldBuffer<BF>,
325	) -> BF {
326		izip!(first_col.as_ref(), second_col.as_ref(), third_col.as_ref(), eq_ind.as_ref())
327			.map(|(a, b, c, eq)| (*a * *b - *c) * *eq)
328			.sum()
329	}
330
331	#[test]
332	fn test_first_round_message_matches_next_round_sum_claim() {
333		// Fixed seed keeps the random witness reproducible across runs.
334		let mut rng = StdRng::from_seed([0; 32]);
335
336		// 2^10 rows total; each 64-bit word packs 2^SKIPPED_VARS rows, leaving this many words.
337		let log_num_words = 10 - SKIPPED_VARS;
338
339		// Every word is random and non-zero, so no window is skipped.
340		// This pins the dense path, where the skip must leave the result unchanged.
341		let mlv_1 = random_words(log_num_words, &mut rng);
342		let mlv_2 = random_words(log_num_words, &mut rng);
343
344		assert_round_message_consistent(&mlv_1, &mlv_2, &mut rng);
345	}
346
347	/// The width in words of one round-message window: `2^(3 + 4) = 128`.
348	const WINDOW: usize = 1 << (PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len() + 4);
349
350	/// The column lengths covering every windowing regime.
351	fn windowing_shapes() -> [usize; 6] {
352		[
353			// A single-row axis: one window, filled by repetition.
354			1,
355			// A sub-window axis with a real zero tail inside it, then repeated.
356			3,
357			// Exactly one window, no padding at all.
358			WINDOW,
359			// One whole window plus a straddling one, padded up to two.
360			WINDOW + 1,
361			// A whole number of windows, padded up to four.
362			3 * WINDOW,
363			// Whole windows plus a straddling one, padded up to four.
364			3 * WINDOW + 17,
365		]
366	}
367
368	// An unpadded column's round message agrees with the verifier's own fold of that column.
369	//
370	// The padded-vs-unpadded equality above would also hold if both were wrong in the same way.
371	// This pins the message to the independently folded claim at every windowing shape.
372	#[test]
373	fn test_first_round_message_with_unpadded_columns() {
374		let mut rng = StdRng::from_seed([3; 32]);
375
376		for n_words in windowing_shapes() {
377			let [a, b] = array::from_fn(|_| {
378				repeat_with(|| Word(rng.random()))
379					.take(n_words)
380					.collect::<Vec<_>>()
381			});
382			assert_round_message_consistent(&a, &b, &mut rng);
383		}
384	}
385
386	/// Asserts the first-round univariate message agrees with the next-round sum claim.
387	///
388	/// The check mirrors the verifier at one random challenge:
389	/// - Extrapolate the round message at the challenge to get the expected next-round sum.
390	/// - Fold A, B, and C = A & B at the same challenge, then form the sum claim directly.
391	/// - The two values must be equal.
392	///
393	/// The columns need not have a power-of-two length; the constraint axis is then the next power
394	/// of two, and both sides read the rows past the columns' end as zero.
395	fn assert_round_message_consistent(mlv_1: &[Word], mlv_2: &[Word], mut rng: impl Rng) {
396		assert_eq!(mlv_1.len(), mlv_2.len());
397		let log_num_words = log2_ceil_usize(mlv_1.len());
398
399		// The prover pins only as many small-field challenges as the axis has coordinates, so an
400		// axis shorter than the fixed set uses a prefix of it.
401		let small_field_zerocheck_challenges = &PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES
402			[..log_num_words.min(PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len())];
403
404		let big_field_zerocheck_challenges =
405			vec![
406				B128::random(&mut rng);
407				log_num_words.saturating_sub(PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len())
408			];
409
410		// The round message derives C = A & B internally.
411		// This materialized copy feeds only the verifier-side transparent fold below.
412		let mlv_3: Vec<Word> = iter::zip(mlv_1, mlv_2).map(|(&a, &b)| a & b).collect();
413
414		// Agreed-upon proof parameter
415
416		let prover_message_domain = BinarySubspace::with_dim(SKIPPED_VARS + 1);
417
418		let verifier_message_domain = prover_message_domain.isomorphic::<B128>();
419
420		// Prover generates first round message
421		let first_round_message_on_ext_domain = univariate_round_message_extension_domain::<B128>(
422			log_num_words,
423			mlv_1,
424			mlv_2,
425			&big_field_zerocheck_challenges,
426			&prover_message_domain,
427		);
428
429		let mut first_round_message_coeffs = vec![B128::ZERO; 2 * ROWS_PER_HYPERCUBE_VERTEX];
430
431		first_round_message_coeffs[ROWS_PER_HYPERCUBE_VERTEX..2 * ROWS_PER_HYPERCUBE_VERTEX]
432			.copy_from_slice(&first_round_message_on_ext_domain);
433
434		// Verifier checks the accuracy of the message by challenging the prover and folding
435		// polynomials transparently
436
437		let verifier_input_domain: BinarySubspace<B128> =
438			verifier_message_domain.reduce_dim(verifier_message_domain.dim() - 1);
439
440		let first_sumcheck_challenge = B128::random(&mut rng);
441		let expected_next_round_sum = verifier_message_domain
442			.extrapolate(&first_round_message_coeffs, &first_sumcheck_challenge);
443
444		let lagrange_evals = verifier_input_domain.lagrange_evals(&first_sumcheck_challenge);
445		let folder = BitAxisFolder::new(&lagrange_evals);
446
447		let folded_first_mle: FieldBuffer<B128> = folder.fold(&GlobalAllocator, mlv_1);
448		let folded_second_mle: FieldBuffer<B128> = folder.fold(&GlobalAllocator, mlv_2);
449		let folded_third_mle: FieldBuffer<B128> = folder.fold(&GlobalAllocator, &mlv_3);
450
451		let upcasted_small_field_challenges: Vec<_> = small_field_zerocheck_challenges
452			.iter()
453			.copied()
454			.map(B128::from)
455			.collect();
456
457		let verifier_field_zerocheck_challenges: Vec<_> = upcasted_small_field_challenges
458			.iter()
459			.chain(big_field_zerocheck_challenges.iter())
460			.copied()
461			.collect();
462
463		let verifier_field_eq = eq_ind_partial_eval(&verifier_field_zerocheck_challenges);
464		let actual_next_round_sum =
465			sum_claim(&folded_first_mle, &folded_second_mle, &folded_third_mle, &verifier_field_eq);
466
467		assert_eq!(expected_next_round_sum, actual_next_round_sum);
468	}
469}