1use std::{array, iter, ops::Deref};
5
6use binius_compute::{Allocator, VecLike};
7use binius_core::word::Word;
8use binius_field::{
9 ExtensionField, PackedField,
10 linear_transformation::{
11 BytewiseLookupTransformationFactory, InputWrappingTransformationFactory,
12 LinearTransformationFactory, OutputWrappingTransformationFactory, Transformation,
13 },
14 packed_extension,
15};
16use binius_ip_prover::{
17 channel::IPProverChannel,
18 sumcheck::{bivariate_product_prover, prove_single},
19};
20use binius_math::{
21 FieldBuffer, FieldSlice, FieldVec, inner_product::inner_product_packed,
22 multilinear::eq::eq_ind_partial_eval, tensor_algebra::TensorAlgebra,
23};
24use binius_utils::{checked_arithmetics::log2_ceil_usize, rayon::prelude::*};
25use binius_verifier::{
26 config::{B1, B128, LOG_WORDS_PER_ELEM},
27 protocols::shift::evaluate_words_mle,
28};
29use itertools::izip;
30
31use crate::{
32 bit_matrix::{ColumnSums, LOG_WEIGHTS_PER_TABLE, RowFoldTables},
33 prove::pack_witness,
34};
35
36pub const LOG_SPLIT_BLOCK: usize = <B128 as ExtensionField<B1>>::LOG_DEGREE;
44
45const N_ROW_TABLES: usize = 1 << (LOG_SPLIT_BLOCK - LOG_WEIGHTS_PER_TABLE);
47
48pub fn fold_1b_rows_for_b128_split<P, Data>(
96 mat: &FieldBuffer<P, Data>,
97 eq_lo: &FieldBuffer<B128>,
98 eq_hi: &FieldBuffer<B128>,
99) -> FieldBuffer<B128>
100where
101 P: PackedField<Scalar = B128>,
102 Data: Deref<Target = [P]>,
103{
104 let log_scalar_bit_width = <B128 as ExtensionField<B1>>::LOG_DEGREE;
105 assert_eq!(mat.log_len(), eq_lo.log_len() + eq_hi.log_len()); assert!(eq_lo.log_len() >= P::LOG_WIDTH || eq_hi.log_len() == 0); assert!(eq_lo.log_len() <= LOG_SPLIT_BLOCK); let lo_tables = RowFoldTables::<B128, N_ROW_TABLES>::new(eq_lo.as_ref());
115
116 let block_packed_len = 1 << eq_lo.log_len().saturating_sub(P::LOG_WIDTH);
118
119 (mat.as_ref().par_chunks(block_packed_len), eq_hi.as_ref().par_iter())
120 .into_par_iter()
121 .fold(
122 || FieldBuffer::zeros(log_scalar_bit_width),
123 |mut acc, (mat_block, &eq_hi_val)| {
124 let mut rows = P::iter_slice(mat_block);
130 let mut sums = ColumnSums::zero();
131
132 lo_tables.fold_into(
139 iter::repeat_with(|| {
140 array::from_fn(|_| {
141 rows.next()
142 .map(packed_extension::cast_base::<B1, _>)
143 .unwrap_or_default()
144 })
145 })
146 .take(N_ROW_TABLES),
147 &mut sums,
148 );
149
150 sums.add_scaled_to(eq_hi_val, acc.as_mut());
153 acc
154 },
155 )
156 .reduce_with(|mut lhs, rhs| {
159 for (lhs_i, &rhs_i) in izip!(lhs.as_mut(), rhs.as_ref()) {
160 *lhs_i += rhs_i;
161 }
162 lhs
163 })
164 .unwrap_or_else(|| FieldBuffer::zeros(log_scalar_bit_width))
166}
167
168pub fn rs_eq_ind_from_factors<A, P>(
197 alloc: &A,
198 eq_lo: &FieldBuffer<B128>,
199 eq_hi: &FieldBuffer<B128>,
200 row_batch_query: &FieldBuffer<B128>,
201) -> FieldVec<P, A>
202where
203 A: Allocator,
204 P: PackedField<Scalar = B128>,
205{
206 assert!(eq_lo.log_len() >= P::LOG_WIDTH || eq_hi.log_len() == 0); assert_eq!(row_batch_query.log_len(), <B128 as ExtensionField<B1>>::LOG_DEGREE); let transform = OutputWrappingTransformationFactory::new(
211 InputWrappingTransformationFactory::new(BytewiseLookupTransformationFactory),
212 )
213 .create(row_batch_query.as_ref());
214
215 let log_len = eq_lo.log_len() + eq_hi.log_len();
216 let packed_len = 1usize << log_len.saturating_sub(P::LOG_WIDTH);
217 let block_packed_len = 1 << eq_lo.log_len().saturating_sub(P::LOG_WIDTH);
218
219 let mut out = alloc.alloc::<P>(packed_len);
221 (
222 out.spare_capacity_mut()[..packed_len].par_chunks_mut(block_packed_len),
223 eq_hi.as_ref().par_iter(),
224 )
225 .into_par_iter()
226 .for_each(|(out_block, &eq_hi_val)| {
227 let lo_chunks = eq_lo.as_ref().chunks(P::WIDTH);
232 for (slot, lo_chunk) in iter::zip(out_block, lo_chunks) {
233 slot.write(P::from_scalars(
234 lo_chunk
235 .iter()
236 .map(|&lo| transform.transform(&(lo * eq_hi_val))),
237 ));
238 }
239 });
240 unsafe { out.set_len(packed_len) };
242
243 FieldBuffer::new(log_len, out)
244}
245
246fn expand_tensor_factors(point: &[B128]) -> (FieldBuffer<B128>, FieldBuffer<B128>) {
251 let (point_lo, point_hi) = point.split_at(point.len().min(LOG_SPLIT_BLOCK));
252 (eq_ind_partial_eval::<B128>(point_lo), eq_ind_partial_eval::<B128>(point_hi))
253}
254
255pub struct RingSwitchOutput<A: Allocator, P: PackedField> {
257 pub rs_eq_ind: FieldVec<P, A>,
259 pub sumcheck_claim: P::Scalar,
261}
262
263pub fn prove<A, P, Channel>(
285 alloc: &A,
286 packed_witness: FieldSlice<'_, P>,
287 eval_point: &[B128],
288 channel: &mut Channel,
289) -> RingSwitchOutput<A, P>
290where
291 A: Allocator,
292 P: PackedField<Scalar = B128>,
293 Channel: IPProverChannel<B128>,
294{
295 let log_packing = <B128 as ExtensionField<B1>>::LOG_DEGREE;
296 assert_eq!(packed_witness.log_len() + log_packing, eval_point.len());
297
298 let eval_point_suffix = &eval_point[log_packing..];
299 let (eq_lo, eq_hi) = tracing::debug_span!("Expand evaluation suffix query")
300 .in_scope(|| expand_tensor_factors(eval_point_suffix));
301
302 let s_hat_v = tracing::debug_span!("Compute ring-switching partial evaluations")
304 .in_scope(|| fold_1b_rows_for_b128_split(&packed_witness, &eq_lo, &eq_hi));
305 channel.send_many(s_hat_v.as_ref());
306
307 let s_hat_u = TensorAlgebra::<B1, B128>::new(s_hat_v.as_ref().to_vec())
309 .transpose()
310 .elems;
311
312 let r_double_prime = channel.sample_many(log_packing);
314 let eq_r_double_prime = eq_ind_partial_eval::<B128>(&r_double_prime);
315
316 let sumcheck_claim = inner_product_packed::<B128, B128>(
319 log_packing,
320 s_hat_u.into_iter(),
321 eq_r_double_prime.as_ref().iter().copied(),
322 );
323
324 let rs_eq_ind = tracing::debug_span!("Compute ring-switching equality indicator")
326 .in_scope(|| rs_eq_ind_from_factors::<A, P>(alloc, &eq_lo, &eq_hi, &eq_r_double_prime));
327
328 RingSwitchOutput {
329 rs_eq_ind,
330 sumcheck_claim,
331 }
332}
333
334pub fn prove_public_eval<A, P, Channel>(
359 alloc: &A,
360 public_words: &[Word],
361 r_j: &[B128],
362 r_y: &[B128],
363 channel: &mut Channel,
364) where
365 A: Allocator,
366 P: PackedField<Scalar = B128>,
367 Channel: IPProverChannel<B128>,
368{
369 let log_public_elems = log2_ceil_usize(public_words.len()).saturating_sub(LOG_WORDS_PER_ELEM);
372 let r_y_public = &r_y[..log_public_elems + LOG_WORDS_PER_ELEM];
373
374 channel.send_one(evaluate_words_mle::<B128, B128>(public_words, r_j, r_y_public));
375
376 let packed = pack_witness::<P, _>(alloc, log_public_elems, public_words)
377 .expect("the element count is derived from the words being packed");
378 let RingSwitchOutput {
379 rs_eq_ind,
380 sumcheck_claim,
381 } = prove(alloc, packed.as_view(), &[r_j, r_y_public].concat(), channel);
382
383 let prover = bivariate_product_prover(alloc, [packed, rs_eq_ind], sumcheck_claim);
387 prove_single(prover, channel);
388}
389
390#[cfg(test)]
391mod test {
392 use binius_compute::GlobalAllocator;
393 use binius_field::{
394 ExtensionField, Field, Ghash128b, PackedField, PackedGhash2x128b, PackedGhash4x128b,
395 PackedSubfield, packed_extension,
396 };
397 use binius_math::{
398 FieldBuffer,
399 inner_product::{inner_product_buffers, inner_product_subfield},
400 multilinear::{eq::eq_ind_partial_eval, evaluate::evaluate_inplace},
401 test_utils::{index_to_hypercube_point, random_field_buffer, random_scalars},
402 };
403 use binius_verifier::{config::B1, ring_switch::eval_rs_eq};
404 use rand::{SeedableRng, rngs::StdRng};
405
406 use super::*;
407
408 type F = Ghash128b;
409
410 fn naive_fold_1b_rows<P: PackedField<Scalar = F>>(mat: &FieldBuffer<P>, eq: &[F]) -> Vec<F> {
417 let mut out = vec![F::ZERO; <F as ExtensionField<B1>>::DEGREE];
418 for (r, &weight) in eq.iter().enumerate() {
419 let row = mat.get(r);
420 for (bit, out_b) in iter::zip(ExtensionField::<B1>::iter_bases(&row), &mut out) {
421 if bit == B1::ONE {
422 *out_b += weight;
423 }
424 }
425 }
426 out
427 }
428
429 fn check_split_fold_matches_definition<P: PackedField<Scalar = F>>(log_len: usize, seed: u64) {
436 let mut rng = StdRng::seed_from_u64(seed);
437 let mat = random_field_buffer::<P>(&mut rng, log_len);
438 let point: Vec<F> = random_scalars(&mut rng, log_len);
439
440 let expected = naive_fold_1b_rows(&mat, eq_ind_partial_eval::<F>(&point).as_ref());
441
442 for split_at in P::LOG_WIDTH.min(log_len)..=log_len.min(LOG_SPLIT_BLOCK) {
445 let (point_lo, point_hi) = point.split_at(split_at);
446 let eq_lo = eq_ind_partial_eval::<F>(point_lo);
447 let eq_hi = eq_ind_partial_eval::<F>(point_hi);
448
449 let split = fold_1b_rows_for_b128_split(&mat, &eq_lo, &eq_hi);
450 assert_eq!(
451 split.as_ref(),
452 expected.as_slice(),
453 "log_len={log_len}, split_at={split_at}"
454 );
455 }
456 }
457
458 #[test]
459 fn test_split_fold_matches_definition() {
460 for (i, log_len) in [0, 1, 2, 6, 7, 8].into_iter().enumerate() {
462 let seed = i as u64;
463 check_split_fold_matches_definition::<F>(log_len, seed);
464 check_split_fold_matches_definition::<PackedGhash2x128b>(log_len, seed);
465 check_split_fold_matches_definition::<PackedGhash4x128b>(log_len, seed);
466 }
467 }
468
469 fn check_rs_eq_ind_from_factors<P: PackedField<Scalar = F>>(log_len: usize, seed: u64) {
476 let mut rng = StdRng::seed_from_u64(seed);
477 let point: Vec<F> = random_scalars(&mut rng, log_len);
478 let row_batching_challenges: Vec<F> =
479 random_scalars(&mut rng, <F as ExtensionField<B1>>::LOG_DEGREE);
480 let row_batch_query = eq_ind_partial_eval::<F>(&row_batching_challenges);
481
482 for split_at in P::LOG_WIDTH.min(log_len)..=log_len {
484 let (point_lo, point_hi) = point.split_at(split_at);
485 let eq_lo = eq_ind_partial_eval::<F>(point_lo);
486 let eq_hi = eq_ind_partial_eval::<F>(point_hi);
487
488 let rs_eq_ind =
489 rs_eq_ind_from_factors::<_, P>(&GlobalAllocator, &eq_lo, &eq_hi, &row_batch_query);
490
491 for index in 0..1 << log_len {
492 let expected = eval_rs_eq::<F>(
493 &point,
494 &index_to_hypercube_point::<F>(log_len, index),
495 row_batch_query.as_ref(),
496 );
497 assert_eq!(rs_eq_ind.get(index), expected, "split_at={split_at}, index={index}");
498 }
499
500 let trailing = rs_eq_ind.as_ref()[0].iter().skip(1 << log_len);
503 assert!(trailing.take(P::WIDTH).all(|lane| lane == F::ZERO), "split_at={split_at}");
504 }
505 }
506
507 #[test]
508 fn test_rs_eq_ind_from_factors() {
509 for (i, log_len) in [0, 1, 2, 6, 7, 8].into_iter().enumerate() {
511 let seed = i as u64;
512 check_rs_eq_ind_from_factors::<F>(log_len, seed);
513 check_rs_eq_ind_from_factors::<PackedGhash2x128b>(log_len, seed);
514 check_rs_eq_ind_from_factors::<PackedGhash4x128b>(log_len, seed);
515 }
516 }
517
518 #[test]
519 fn test_out_of_range_evaluation() {
520 let mut rng = StdRng::from_seed([0; 32]);
521
522 let n_vars_big_field = 3;
526
527 let z_vals: Vec<F> = random_scalars(&mut rng, n_vars_big_field);
529
530 let row_batching_challenges: Vec<F> =
531 random_scalars(&mut rng, <F as ExtensionField<B1>>::LOG_DEGREE);
532
533 let row_batching_expanded_query: FieldBuffer<F> =
534 eq_ind_partial_eval(&row_batching_challenges);
535
536 let (eq_lo, eq_hi) = expand_tensor_factors(&z_vals);
537 let rs_eq = rs_eq_ind_from_factors::<_, F>(
538 &GlobalAllocator,
539 &eq_lo,
540 &eq_hi,
541 &row_batching_expanded_query,
542 );
543
544 let eval_point: Vec<F> = random_scalars(&mut rng, n_vars_big_field);
546
547 let tensor_expanded_eval_point = eq_ind_partial_eval::<F>(&eval_point);
550 let expected_eval = inner_product_buffers(&rs_eq, &tensor_expanded_eval_point);
551
552 let actual_eval =
553 eval_rs_eq::<F>(&z_vals, &eval_point, row_batching_expanded_query.as_ref());
554
555 assert_eq!(expected_eval, actual_eval);
556 }
557
558 #[test]
559 fn test_row_fold_composes_into_the_claim() {
560 let mut rng = StdRng::seed_from_u64(0);
561
562 type P = PackedGhash2x128b;
563
564 let n = 10;
571 let log_degree = <F as ExtensionField<B1>>::LOG_DEGREE;
572
573 let bit_matrix = random_field_buffer::<PackedSubfield<P, B1>>(&mut rng, n + log_degree);
575 let eval_point: Vec<F> = random_scalars(&mut rng, n + log_degree);
576 let (prefix, suffix) = eval_point.split_at(log_degree);
577
578 let full_tensor = eq_ind_partial_eval::<F>(&eval_point);
580 let expected = inner_product_subfield(
581 PackedField::iter_slice(bit_matrix.as_ref()),
582 PackedField::iter_slice(full_tensor.as_ref()),
583 );
584
585 let mat = FieldBuffer::<P>::new(
589 n,
590 bit_matrix
591 .as_ref()
592 .iter()
593 .map(|&bits_packed| packed_extension::cast_ext::<B1, P>(bits_packed))
594 .collect(),
595 );
596 let (eq_lo, eq_hi) = expand_tensor_factors(suffix);
597 let s_hat_v = fold_1b_rows_for_b128_split(&mat, &eq_lo, &eq_hi);
598
599 assert_eq!(evaluate_inplace(s_hat_v, prefix), expected);
600 }
601}