binius_prover/protocols/bitand/
sumcheck_round_messages.rs1use 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
20const N_FIXED_LARGE_CHALLENGES: usize = 4;
35
36pub 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 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 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 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
189fn 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
223fn 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 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
273struct 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 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 let mut rng = StdRng::from_seed([0; 32]);
335
336 let log_num_words = 10 - SKIPPED_VARS;
338
339 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 const WINDOW: usize = 1 << (PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES.len() + 4);
349
350 fn windowing_shapes() -> [usize; 6] {
352 [
353 1,
355 3,
357 WINDOW,
359 WINDOW + 1,
361 3 * WINDOW,
363 3 * WINDOW + 17,
365 ]
366 }
367
368 #[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 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 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 let mlv_3: Vec<Word> = iter::zip(mlv_1, mlv_2).map(|(&a, &b)| a & b).collect();
413
414 let prover_message_domain = BinarySubspace::with_dim(SKIPPED_VARS + 1);
417
418 let verifier_message_domain = prover_message_domain.isomorphic::<B128>();
419
420 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 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}