1use std::{iter, ops::Deref};
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, PackedField, WideMul};
9use binius_ip::sumcheck::RoundCoeffs;
10use binius_ip_prover::{
11 channel::IPProverChannel,
12 sumcheck::{
13 ProveSingleOutput, bivariate_product_evaluator::BivariateProductEvaluator,
14 bivariate_product_prover, common::SumcheckProver, prove_single, round_evals::RoundEvals,
15 round_evaluator::SharedSumcheckProver,
16 },
17};
18use binius_math::{FieldBuffer, FieldVec, multilinear::fold::fold_highest_var_inplace};
19use binius_verifier::protocols::shift::LOG_SHIFT_COUNT;
20use tracing::instrument;
21
22use super::{
23 SegmentWords,
24 claims::PreparedOperandClaims,
25 key_collection::{DenseShiftEncoding, KeyCollection},
26 monster::shift_operator_table,
27 outer::OuterShiftStage,
28 shift_ind::ShiftChallenge,
29};
30
31#[instrument(skip_all, name = "prover_phase_1")]
50pub fn prove_phase_1<F, P, Channel, A>(
51 key_collection: &KeyCollection,
52 words: SegmentWords<'_>,
53 prepared: &PreparedOperandClaims<F>,
54 oblong_weights: &[F],
55 channel: &mut Channel,
56 alloc: &A,
57) -> Phase1Output<F>
58where
59 F: BinaryField,
60 P: PackedField<Scalar = F>,
61 Channel: IPProverChannel<F>,
62 A: Allocator,
63{
64 let public = key_collection
68 .public
69 .build_g::<_, P>(words.public, prepared);
70 let hidden = key_collection
71 .hidden
72 .build_g::<_, P>(words.hidden, prepared);
73 let g = SparseShiftRows::from_segments([
74 (&public, &key_collection.public.dense_shift_enc),
75 (&hidden, &key_collection.hidden.dense_shift_enc),
76 ]);
77
78 g.run_phase_1_sumcheck(oblong_weights, prepared.batched_eval, channel, alloc)
79}
80
81pub const PHASE_1_LOG_LEN: usize = Word::LOG_BITS + LOG_SHIFT_ROWS;
87
88pub const LOG_SHIFT_ROWS: usize = 2 * LOG_SHIFT_COUNT;
94
95pub const SHIFT_OPERATOR_LOG_LEN: usize = Word::LOG_BITS + LOG_SHIFT_COUNT;
98
99#[derive(Debug, Clone)]
104pub struct Phase1Output<F> {
105 pub r_j: Vec<F>,
107 pub inner: ShiftChallenge<F>,
109 pub outer: ShiftChallenge<F>,
111 pub psi: Vec<F>,
115 pub gamma: F,
117 pub g_eval: F,
121}
122
123pub(super) const fn row_len<P: PackedField>() -> usize {
125 assert!(
126 P::LOG_WIDTH <= Word::LOG_BITS,
127 "a row of `Word::BITS` scalars must be a whole number of packed elements"
128 );
129 Word::BITS >> P::LOG_WIDTH
130}
131
132#[derive(Debug, Clone)]
142pub struct SparseShiftRows<P: PackedField> {
143 indices: Vec<u32>,
147 values: Vec<P>,
149 log_rows: usize,
151}
152
153impl<P: PackedField> SparseShiftRows<P> {
154 pub fn from_segments(segments: [(&[P], &DenseShiftEncoding); 2]) -> Self {
165 let mut indices = Vec::new();
166 let mut values = Vec::new();
167
168 for (rows, dense_shift_enc) in segments {
170 assert_eq!(
171 rows.len(),
172 dense_shift_enc.len() * row_len::<P>(),
173 "a segment holds one row per shift its encoding names"
174 );
175 indices.extend(dense_shift_enc.shift_indices().map(|index| index as u32));
177 values.extend_from_slice(rows);
179 }
180
181 Self::new(indices, values, LOG_SHIFT_ROWS)
182 }
183
184 pub fn new(indices: Vec<u32>, values: Vec<P>, log_rows: usize) -> Self {
194 assert_eq!(
195 values.len(),
196 indices.len() * row_len::<P>(),
197 "the values hold one row per index"
198 );
199 assert!(
200 indices
201 .iter()
202 .all(|&index| (index as usize) < 1 << log_rows),
203 "every index names a row of the space"
204 );
205
206 Self {
207 indices,
208 values,
209 log_rows,
210 }
211 }
212
213 pub(crate) const fn log_rows(&self) -> usize {
215 self.log_rows
216 }
217
218 pub(crate) fn rows(&self) -> impl Iterator<Item = (usize, &[P])> {
220 iter::zip(&self.indices, self.values.chunks_exact(row_len::<P>()))
221 .map(|(&index, row)| (index as usize, row))
222 }
223
224 pub(crate) fn half(&self) -> usize {
230 assert!(self.log_rows > 0, "precondition: a row-index variable remains to bind");
231 1 << (self.log_rows - 1)
232 }
233
234 pub(crate) fn fold(&mut self, challenge: P::Scalar) {
244 let half = self.half();
245 let lower_weight = P::broadcast(P::Scalar::ONE - challenge);
246 let upper_weight = P::broadcast(challenge);
247
248 let row_len = row_len::<P>();
249 for (position, index) in self.indices.iter_mut().enumerate() {
252 let row = &mut self.values[position * row_len..][..row_len];
253 if *index as usize & half == 0 {
254 row.iter_mut().for_each(|value| *value *= lower_weight);
256 } else {
257 row.iter_mut().for_each(|value| *value *= upper_weight);
259 *index ^= half as u32;
260 }
261 }
262
263 self.log_rows -= 1;
264 }
265
266 fn into_bit_multilinear<A: Allocator>(self, alloc: &A) -> FieldVec<P, A> {
273 assert_eq!(self.log_rows, 0, "precondition: every row-index variable is bound");
274
275 let mut multilinear = FieldBuffer::zeros_in(alloc, Word::LOG_BITS);
276 for (_, row) in self.rows() {
278 for (slot, &value) in iter::zip(multilinear.as_mut(), row) {
279 *slot += value;
280 }
281 }
282 multilinear
283 }
284
285 fn round_coeffs<F, Data>(&self, h: &FieldBuffer<P, Data>, claim: F) -> RoundCoeffs<F>
296 where
297 F: Field,
298 P: PackedField<Scalar = F>,
299 Data: Deref<Target = [P]>,
300 {
301 let half = self.half();
302 let row_len = row_len::<P>();
303 let h_rows = h.as_ref();
304
305 let mut y_1 = <P as WideMul>::Output::default();
308 let mut y_inf = <P as WideMul>::Output::default();
309 for (index, row) in self.rows() {
310 let own = &h_rows[index * row_len..][..row_len];
311 let facing = &h_rows[(index ^ half) * row_len..][..row_len];
312
313 for i in 0..row_len {
314 if index & half != 0 {
317 y_1 += P::wide_mul(row[i], own[i]);
318 }
319 y_inf += P::wide_mul(row[i], own[i] + facing[i]);
320 }
321 }
322
323 let sum_lanes = |wide| P::reduce(wide).iter().sum::<F>();
326 RoundEvals([sum_lanes(y_1), sum_lanes(y_inf)]).interpolate(claim)
327 }
328
329 #[instrument(skip_all, name = "run_sumcheck")]
358 pub fn run_phase_1_sumcheck<F, Channel, A>(
359 mut self,
360 oblong_weights: &[F],
361 sum: F,
362 channel: &mut Channel,
363 alloc: &A,
364 ) -> Phase1Output<F>
365 where
366 F: BinaryField,
367 P: PackedField<Scalar = F>,
368 Channel: IPProverChannel<F>,
369 A: Allocator,
370 {
371 assert_eq!(self.log_rows(), LOG_SHIFT_ROWS, "the row list spans both shift slots");
372
373 let mut outer = OuterShiftStage::new(alloc, oblong_weights);
378 let mut claim = sum;
379 let mut outer_point = Vec::with_capacity(LOG_SHIFT_COUNT);
380 for _ in 0..LOG_SHIFT_COUNT {
381 let round_coeffs = outer.round_coeffs(&self, claim);
382 channel.send_many(round_coeffs.clone().truncate().coeffs());
383 let challenge = channel.sample();
384 claim = round_coeffs.evaluate(&challenge);
385 outer.fold(challenge);
386 self.fold(challenge);
387 outer_point.push(challenge);
388 }
389 let psi = outer.psi().to_vec();
391
392 let h = shift_operator_table(alloc, &psi);
397
398 let g = self;
400 let ProveSingleOutput {
401 multilinear_evals,
402 mut challenges,
403 } = prove_single(Phase1SumcheckProver::new(g, h, claim, alloc), channel);
404
405 challenges.reverse();
409 assert_eq!(challenges.len(), SHIFT_OPERATOR_LOG_LEN);
410 let mut r_j = challenges;
411 let r_v_inner = r_j.split_off(Word::LOG_BITS * 2);
412 let r_s_inner = r_j.split_off(Word::LOG_BITS);
413
414 outer_point.reverse();
415 let mut r_s_outer = outer_point;
416 let r_v_outer = r_s_outer.split_off(Word::LOG_BITS);
417
418 let [g_eval, h_eval] = multilinear_evals
419 .try_into()
420 .expect("prover has 2 multilinear polynomials");
421
422 Phase1Output {
423 r_j,
424 inner: ShiftChallenge::new(r_s_inner, r_v_inner),
425 outer: ShiftChallenge::new(r_s_outer, r_v_outer),
426 psi,
427 gamma: g_eval * h_eval,
428 g_eval,
429 }
430 }
431}
432
433pub struct Phase1SumcheckProver<'alloc, A: Allocator, P: PackedField> {
453 alloc: &'alloc A,
454 stage: Option<Stage<'alloc, A, P>>,
459}
460
461enum Stage<'alloc, A: Allocator, P: PackedField> {
463 Rows {
465 g: SparseShiftRows<P>,
466 h: FieldVec<P, A>,
467 claim: P::Scalar,
469 coeffs: Option<RoundCoeffs<P::Scalar>>,
474 },
475 Bits(SharedSumcheckProver<'alloc, A, P, BivariateProductEvaluator>),
477}
478
479impl<'alloc, A: Allocator, F: Field, P: PackedField<Scalar = F>>
480 Phase1SumcheckProver<'alloc, A, P>
481{
482 pub fn new(g: SparseShiftRows<P>, h: FieldVec<P, A>, sum: F, alloc: &'alloc A) -> Self {
493 assert_eq!(h.log_len(), SHIFT_OPERATOR_LOG_LEN, "h spans one shift slot");
494 assert_eq!(g.log_rows(), LOG_SHIFT_COUNT, "g's rows span one shift slot");
495
496 Self {
497 alloc,
498 stage: Some(Stage::Rows {
499 g,
500 h,
501 claim: sum,
502 coeffs: None,
503 }),
504 }
505 }
506}
507
508impl<A: Allocator, F: Field, P: PackedField<Scalar = F>> SumcheckProver<F>
509 for Phase1SumcheckProver<'_, A, P>
510{
511 fn n_vars(&self) -> usize {
512 match self.stage.as_ref().expect("the stage is set between calls") {
513 Stage::Rows { h, .. } => h.log_len(),
516 Stage::Bits(prover) => prover.n_vars(),
517 }
518 }
519
520 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
521 match self.stage.as_mut().expect("the stage is set between calls") {
522 Stage::Rows {
523 g,
524 h,
525 claim,
526 coeffs,
527 } => {
528 let round_coeffs = g.round_coeffs(h, *claim);
531 *coeffs = Some(round_coeffs.clone());
532 vec![round_coeffs]
533 }
534 Stage::Bits(prover) => prover.execute(),
536 }
537 }
538
539 fn fold(&mut self, challenge: F) {
540 let stage = self.stage.take().expect("the stage is set between calls");
542 self.stage = Some(match stage {
543 Stage::Rows {
544 mut g,
545 mut h,
546 coeffs,
547 ..
548 } => {
549 let claim = coeffs
550 .expect("execute is called before fold")
551 .evaluate(&challenge);
552 g.fold(challenge);
556 fold_highest_var_inplace(&mut h, challenge);
557
558 if g.log_rows() > 0 {
559 Stage::Rows {
561 g,
562 h,
563 claim,
564 coeffs: None,
565 }
566 } else {
567 Stage::Bits(bivariate_product_prover(
571 self.alloc,
572 [g.into_bit_multilinear(self.alloc), h],
573 claim,
574 ))
575 }
576 }
577 Stage::Bits(mut prover) => {
578 prover.fold(challenge);
580 Stage::Bits(prover)
581 }
582 });
583 }
584
585 fn finish(self) -> Vec<F> {
586 match self.stage.expect("the stage is set between calls") {
587 Stage::Rows { .. } => panic!("finish called before the row index was bound"),
588 Stage::Bits(prover) => prover.finish(),
590 }
591 }
592}
593
594#[cfg(test)]
595mod tests {
596 use binius_compute::GlobalAllocator;
597 use binius_core::constraint_system::{
598 AndConstraint, ConstraintSystem, InoutSegment, Shift, ShiftedValueIndex, ValueIndex,
599 };
600 use binius_field::{Field, Ghash128b, PackedGhash2x128b};
601 use binius_math::{inner_product::inner_product_buffers, test_utils::random_scalars};
602 use binius_transcript::ProverTranscript;
603 use binius_verifier::config::StdChallenger;
604 use rand::{SeedableRng, rngs::StdRng};
605
606 use super::*;
607 use crate::protocols::shift::KeyCollection;
608
609 type F = Ghash128b;
610
611 impl<P: PackedField> SparseShiftRows<P> {
612 fn scatter<A: Allocator>(&self, alloc: &A) -> FieldVec<P, A> {
621 let row_len = row_len::<P>();
622 let mut g = FieldBuffer::zeros_in(alloc, self.log_rows + Word::LOG_BITS);
623 for (index, row) in self.rows() {
624 let slots = &mut g.as_mut()[index * row_len..][..row_len];
626 for (slot, &value) in iter::zip(slots, row) {
627 *slot += value;
628 }
629 }
630 g
631 }
632 }
633
634 fn overlapping_shift_system() -> ConstraintSystem {
639 let public = ValueIndex::constant(1);
640 let hidden = ValueIndex::private(1);
641 ConstraintSystem {
642 constants: vec![Word::ZERO; 4],
643 n_inout: 0,
644 n_private: 4,
645 zero_constraints: Vec::new(),
646 and_constraints: vec![AndConstraint([
647 vec![
648 ShiftedValueIndex::plain(public),
649 ShiftedValueIndex::srl(public, 3),
650 ],
651 vec![ShiftedValueIndex::sar(hidden, 7)],
652 vec![
653 ShiftedValueIndex::rotr(hidden, 1),
654 ShiftedValueIndex::plain(hidden),
655 ],
656 ])],
657 imul_constraints: Vec::new(),
658 bmul_constraints: Vec::new(),
659 }
660 }
661
662 #[test]
667 fn from_segments_concatenates_the_two_encodings() {
668 let key_collection =
669 KeyCollection::build(&overlapping_shift_system(), InoutSegment::Public);
670
671 let segment_rows = |enc: &DenseShiftEncoding, base: u128| {
674 (0..enc.len() * Word::BITS)
675 .map(|i| F::new(base + (i / Word::BITS) as u128))
676 .collect::<Vec<F>>()
677 };
678 let public = segment_rows(&key_collection.public.dense_shift_enc, 0x100);
679 let hidden = segment_rows(&key_collection.hidden.dense_shift_enc, 0x200);
680
681 let g = SparseShiftRows::from_segments([
682 (&public, &key_collection.public.dense_shift_enc),
683 (&hidden, &key_collection.hidden.dense_shift_enc),
684 ]);
685
686 let row_index = |shift: Shift| shift.index() as u32;
690 assert_eq!(
691 g.indices,
692 [
693 row_index(Shift::IDENTITY),
694 row_index(Shift::srl(3)),
695 row_index(Shift::IDENTITY),
696 row_index(Shift::sar(7)),
697 row_index(Shift::rotr(1)),
698 ]
699 );
700
701 let row = |position: usize| &g.values[position * Word::BITS..][..Word::BITS];
703 for (position, expected) in [0x100, 0x101, 0x200, 0x201, 0x202].into_iter().enumerate() {
704 assert!(row(position).iter().all(|&value| value == F::new(expected)));
705 }
706
707 let at = |shift: Shift| {
709 g.rows()
710 .filter(|&(index, _)| index == shift.index())
711 .map(|(_, row)| row[0])
712 .sum::<F>()
713 };
714 assert_eq!(at(Shift::IDENTITY), F::new(0x100) + F::new(0x200));
715 assert_eq!(at(Shift::srl(3)), F::new(0x101));
716 assert_eq!(at(Shift::sar(7)), F::new(0x201));
717 assert_eq!(at(Shift::rotr(1)), F::new(0x202));
718 }
719
720 #[test]
722 fn a_sequence_is_keyed_outer_major() {
723 let hidden = ValueIndex::private(1);
724 let sequence = [Shift::srl(3), Shift::sll(5)];
725 let cs = ConstraintSystem {
726 constants: vec![Word::ZERO; 4],
727 n_inout: 0,
728 n_private: 4,
729 zero_constraints: Vec::new(),
730 and_constraints: vec![AndConstraint([
731 vec![ShiftedValueIndex::new(hidden, sequence)],
732 Vec::new(),
733 Vec::new(),
734 ])],
735 imul_constraints: Vec::new(),
736 bmul_constraints: Vec::new(),
737 };
738
739 let key_collection = KeyCollection::build(&cs, InoutSegment::Public);
740 let [inner, outer] = sequence;
741 assert_eq!(
742 key_collection
743 .hidden
744 .dense_shift_enc
745 .shift_indices()
746 .collect::<Vec<_>>(),
747 [outer.index() << LOG_SHIFT_COUNT | inner.index()]
748 );
749 }
750
751 #[test]
756 fn scatter_places_rows_at_their_shift_index() {
757 let indices = [Shift::IDENTITY, Shift::sar(7), Shift::rotr(1)]
758 .map(|shift| shift.index() as u32)
759 .to_vec();
760 let values = (0..indices.len() * Word::BITS)
761 .map(|i| F::new(1 + (i / Word::BITS) as u128))
762 .collect::<Vec<F>>();
763 let rows = SparseShiftRows::<F>::new(indices.clone(), values, LOG_SHIFT_COUNT);
764
765 let g = rows.scatter(&GlobalAllocator);
766 assert_eq!(g.log_len(), SHIFT_OPERATOR_LOG_LEN);
767
768 for (row, &shift_index) in indices.iter().enumerate() {
770 let offset = shift_index as usize * Word::BITS;
771 for bit in 0..Word::BITS {
772 assert_eq!(g.get(offset + bit), F::new(1 + row as u128));
773 }
774 }
775 for row in (0..1 << LOG_SHIFT_COUNT).filter(|row| !indices.contains(&(*row as u32))) {
776 assert!((0..Word::BITS).all(|bit| g.get(row * Word::BITS + bit) == F::ZERO));
777 }
778 }
779
780 fn run_dense_reference<P: PackedField<Scalar = F>>(
786 g: &SparseShiftRows<P>,
787 h: FieldVec<P, GlobalAllocator>,
788 sum: F,
789 channel: &mut ProverTranscript<StdChallenger>,
790 ) -> (Vec<F>, F) {
791 let prover =
792 bivariate_product_prover(&GlobalAllocator, [g.scatter(&GlobalAllocator), h], sum);
793
794 let ProveSingleOutput {
795 multilinear_evals,
796 mut challenges,
797 } = prove_single(prover, channel);
798 challenges.reverse();
799
800 let [g_eval, h_eval] = multilinear_evals
801 .try_into()
802 .expect("prover has 2 multilinear polynomials");
803
804 (challenges, g_eval * h_eval)
805 }
806
807 fn phase_1_multilinears<P: PackedField<Scalar = F>>(
814 cs: &ConstraintSystem,
815 seed: u64,
816 ) -> (SparseShiftRows<P>, FieldVec<P, GlobalAllocator>) {
817 let mut rng = StdRng::seed_from_u64(seed);
818 let key_collection = KeyCollection::build(cs, InoutSegment::Public);
819
820 let mut indices = Vec::new();
821 let mut values = Vec::new();
822 for segment in [&key_collection.public, &key_collection.hidden] {
823 for [inner, _] in segment.dense_shift_enc.iter() {
824 indices.push(inner.index() as u32);
825 values.extend((0..row_len::<P>()).map(|_| P::random(&mut rng)));
826 }
827 }
828 let g = SparseShiftRows::new(indices, values, LOG_SHIFT_COUNT);
829
830 let h = shift_operator_table(&GlobalAllocator, &random_scalars::<F>(&mut rng, Word::BITS));
832
833 (g, h)
834 }
835
836 fn assert_sparse_matches_dense<P: PackedField<Scalar = F>>(cs: &ConstraintSystem, seed: u64) {
842 let (g, h) = phase_1_multilinears::<P>(cs, seed);
843 let sum = inner_product_buffers(&g.scatter(&GlobalAllocator), &h);
845
846 let mut sparse_transcript = ProverTranscript::<StdChallenger>::default();
847 let ProveSingleOutput {
848 multilinear_evals,
849 challenges: mut sparse_challenges,
850 } = prove_single(
851 Phase1SumcheckProver::new(g.clone(), h.clone(), sum, &GlobalAllocator),
852 &mut sparse_transcript,
853 );
854 sparse_challenges.reverse();
855 let [g_eval, h_eval] = multilinear_evals
856 .try_into()
857 .expect("prover has 2 multilinear polynomials");
858
859 let mut dense_transcript = ProverTranscript::<StdChallenger>::default();
860 let (dense_challenges, dense_eval) = run_dense_reference(&g, h, sum, &mut dense_transcript);
861
862 assert_eq!(sparse_challenges, dense_challenges);
863 assert_eq!(g_eval * h_eval, dense_eval);
864 assert_eq!(sparse_transcript.finalize(), dense_transcript.finalize());
865 }
866
867 #[test]
868 fn sparse_rounds_match_the_dense_prover() {
869 assert_sparse_matches_dense::<F>(&overlapping_shift_system(), 0);
872 assert_sparse_matches_dense::<PackedGhash2x128b>(&overlapping_shift_system(), 1);
873 }
874
875 #[test]
880 fn sparse_rounds_match_the_dense_prover_with_an_empty_g() {
881 let cs = ConstraintSystem {
882 constants: vec![Word::ZERO; 4],
883 n_inout: 0,
884 n_private: 4,
885 zero_constraints: Vec::new(),
886 and_constraints: Vec::new(),
887 imul_constraints: Vec::new(),
888 bmul_constraints: Vec::new(),
889 };
890
891 let key_collection = KeyCollection::build(&cs, InoutSegment::Public);
892 assert!(key_collection.public.dense_shift_enc.is_empty());
893 assert!(key_collection.hidden.dense_shift_enc.is_empty());
894
895 assert_sparse_matches_dense::<PackedGhash2x128b>(&cs, 2);
896 }
897}