1use binius_compute::BufferData;
6use binius_field::{Field, PackedField};
7use binius_ip::sumcheck::RoundCoeffs;
8use binius_utils::rayon::{
9 prelude::*,
10 task_size::{IndexedParallelIteratorExt, WorkPerItem},
11};
12
13use super::{
14 common::SumcheckProver, factored_multilinear::FactoredMultilinear, round_evals::RoundEvals,
15 round_state::RoundState,
16};
17
18pub type SparseEntry<F> = (usize, F);
22
23pub struct SparseDenseProductSumcheckProver<P: PackedField, Data: BufferData<P> = Vec<P>> {
68 sparse: Vec<SparseEntry<P::Scalar>>,
72 dense: FactoredMultilinear<P, Data>,
79 state: RoundState<RoundCoeffs<P::Scalar>, P::Scalar>,
81}
82
83impl<P: PackedField, Data: BufferData<P>> SparseDenseProductSumcheckProver<P, Data> {
84 pub fn new(
96 sparse: Vec<SparseEntry<P::Scalar>>,
97 dense: FactoredMultilinear<P, Data>,
98 sum: P::Scalar,
99 ) -> Self {
100 assert!(
101 sparse.iter().all(|&(index, _)| index < 1 << dense.n_vars()),
102 "precondition: every sparse index must be within the dense multilinear"
103 );
104
105 Self {
106 sparse,
107 dense,
108 state: RoundState::Claim(sum),
109 }
110 }
111
112 fn half(&self) -> usize {
118 let n_vars = self.dense.n_vars();
119 assert!(n_vars > 0, "no variables remain to bind");
120 1 << (n_vars - 1)
121 }
122}
123
124impl<F: Field, P: PackedField<Scalar = F>, Data: BufferData<P> + Sync> SumcheckProver<F>
125 for SparseDenseProductSumcheckProver<P, Data>
126{
127 fn n_vars(&self) -> usize {
128 self.dense.n_vars()
129 }
130
131 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
132 let claim = *self.state.claim();
133 let half = self.half();
134
135 let (y_1, y_inf) = self
139 .sparse
140 .par_iter()
141 .with_min_task(WorkPerItem::FieldMuls)
142 .map(|&(index, value)| {
143 let own = self.dense.get(index);
144 let facing = self.dense.get(index ^ half);
145
146 let y_1 = if index & half == 0 {
149 F::ZERO
150 } else {
151 value * own
152 };
153 (y_1, value * (own + facing))
154 })
155 .reduce(
156 || (F::ZERO, F::ZERO),
157 |(lhs_1, lhs_inf), (rhs_1, rhs_inf)| (lhs_1 + rhs_1, lhs_inf + rhs_inf),
158 );
159
160 let coeffs = RoundEvals([y_1, y_inf]).interpolate(claim);
161 self.state = RoundState::Coeffs(coeffs.clone());
162 vec![coeffs]
163 }
164
165 fn fold(&mut self, challenge: F) {
166 let claim = self.state.coeffs().evaluate(&challenge);
167 let half = self.half();
168
169 let lower_weight = F::ONE - challenge;
172 self.sparse
173 .par_iter_mut()
174 .with_min_task(WorkPerItem::FieldMuls)
175 .for_each(|(index, value)| {
176 if *index & half == 0 {
177 *value *= lower_weight;
178 } else {
179 *value *= challenge;
180 *index ^= half;
181 }
182 });
183
184 self.dense.fold_highest_var(challenge);
185 self.state = RoundState::Claim(claim);
186 }
187
188 fn finish(self) -> Vec<F> {
189 assert_eq!(self.n_vars(), 0, "finish called before the last fold");
190
191 let sparse_eval = self.sparse.iter().map(|&(_, value)| value).sum();
194 vec![sparse_eval, self.dense.get(0)]
195 }
196}
197
198pub struct SparseMultiDenseProductSumcheckProver<P: PackedField> {
216 sparse: Vec<SparseEntry<P::Scalar>>,
220
221 dense: Vec<FactoredMultilinear<P>>,
224
225 state: Vec<RoundState<RoundCoeffs<P::Scalar>, P::Scalar>>,
227}
228
229impl<P: PackedField> SparseMultiDenseProductSumcheckProver<P> {
230 pub fn new(
244 sparse: Vec<SparseEntry<P::Scalar>>,
245 dense: Vec<FactoredMultilinear<P>>,
246 sums: &[P::Scalar],
247 ) -> Self {
248 let n_vars = dense
249 .first()
250 .expect("precondition: at least one dense column")
251 .n_vars();
252 assert!(
253 dense.iter().all(|column| column.n_vars() == n_vars),
254 "precondition: every dense column must span the same variables"
255 );
256 assert_eq!(dense.len(), sums.len(), "precondition: one sum per dense column");
257 assert!(
258 sparse.iter().all(|&(index, _)| index < 1 << n_vars),
259 "precondition: every sparse index must be within the dense multilinears"
260 );
261
262 Self {
263 sparse,
264 dense,
265 state: sums.iter().map(|&sum| RoundState::Claim(sum)).collect(),
266 }
267 }
268
269 fn half(&self) -> usize {
275 let n_vars = self.n_vars();
276 assert!(n_vars > 0, "no variables remain to bind");
277 1 << (n_vars - 1)
278 }
279}
280
281impl<F: Field, P: PackedField<Scalar = F>> SumcheckProver<F>
282 for SparseMultiDenseProductSumcheckProver<P>
283{
284 fn n_vars(&self) -> usize {
285 self.dense.first().map_or(0, FactoredMultilinear::n_vars)
286 }
287
288 fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
289 let half = self.half();
290 let sparse = &self.sparse;
291
292 let coeffs = self
299 .dense
300 .iter()
301 .zip(&self.state)
302 .map(|(dense, state)| {
303 let (y_1, y_inf) = sparse
304 .par_iter()
305 .with_min_task(WorkPerItem::FieldMuls)
306 .map(|&(index, value)| {
307 let own = dense.get(index);
308 let facing = dense.get(index ^ half);
309
310 let y_1 = if index & half == 0 {
315 F::ZERO
316 } else {
317 value * own
318 };
319 (y_1, value * (own + facing))
320 })
321 .reduce(
322 || (F::ZERO, F::ZERO),
323 |(lhs_1, lhs_inf), (rhs_1, rhs_inf)| (lhs_1 + rhs_1, lhs_inf + rhs_inf),
324 );
325 RoundEvals([y_1, y_inf]).interpolate(*state.claim())
326 })
327 .collect::<Vec<_>>();
328
329 self.state = coeffs.iter().cloned().map(RoundState::Coeffs).collect();
330 coeffs
331 }
332
333 fn fold(&mut self, challenge: F) {
334 let half = self.half();
335
336 self.state = self
338 .state
339 .iter()
340 .map(|state| RoundState::Claim(state.coeffs().evaluate(&challenge)))
341 .collect();
342
343 let lower_weight = F::ONE - challenge;
345 self.sparse
346 .par_iter_mut()
347 .with_min_task(WorkPerItem::FieldMuls)
348 .for_each(|(index, value)| {
349 if *index & half == 0 {
350 *value *= lower_weight;
351 } else {
352 *value *= challenge;
353 *index ^= half;
354 }
355 });
356
357 for dense in &mut self.dense {
358 dense.fold_highest_var(challenge);
359 }
360 }
361
362 fn finish(self) -> Vec<F> {
363 assert_eq!(self.n_vars(), 0, "finish called before the last fold");
364
365 let sparse_eval = self.sparse.iter().map(|&(_, value)| value).sum();
370 let mut evals = vec![sparse_eval];
371 evals.extend(self.dense.iter().map(|dense| dense.get(0)));
372 evals
373 }
374}
375
376#[cfg(test)]
377mod tests {
378 use binius_compute::GlobalAllocator;
379 use binius_field::{
380 Random,
381 arch::{OptimalB128, OptimalPackedB128},
382 };
383 use binius_ip::sumcheck::verify;
384 use binius_math::{
385 FieldBuffer, multilinear::evaluate::evaluate, test_utils::random_field_buffer,
386 };
387 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
388 use proptest::prelude::*;
389 use rand::{SeedableRng, prelude::*};
390
391 use super::*;
392 use crate::sumcheck::{
393 batch::batch_prove, factored_multilinear::FactoredMultilinear, prove::prove_single,
394 };
395
396 type F = OptimalB128;
397 type P = OptimalPackedB128;
398 type StdChallenger = HasherChallenger<sha2::Sha256>;
399
400 fn densify(sparse: &[SparseEntry<F>], n_vars: usize) -> FieldBuffer<P> {
402 let mut buffer = FieldBuffer::<P>::zeros(n_vars);
403 for &(index, value) in sparse {
404 buffer.set(index, buffer.get(index) + value);
405 }
406 buffer
407 }
408
409 fn prove_verify(sparse: Vec<SparseEntry<F>>, dense: &FieldBuffer<P>) {
412 let n_vars = dense.log_len();
413 let sparse_dense = densify(&sparse, n_vars);
414 let sum = sparse
415 .iter()
416 .map(|&(index, value)| value * dense.get(index))
417 .sum::<F>();
418
419 let weight = FactoredMultilinear::new([dense.clone()]);
421 let prover = SparseDenseProductSumcheckProver::new(sparse, weight, sum);
422
423 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
424 let output = prove_single(prover, &mut prover_transcript);
425 prover_transcript
426 .message()
427 .write_slice(&output.multilinear_evals);
428
429 let mut verifier_transcript = prover_transcript.into_verifier();
430 let sumcheck_output = verify(n_vars, 2, sum, &mut verifier_transcript).unwrap();
431 let multilinear_evals = verifier_transcript.message().read_vec::<F>(2).unwrap();
432
433 assert_eq!(
434 multilinear_evals[0] * multilinear_evals[1],
435 sumcheck_output.eval,
436 "product of the multilinear evaluations should equal the reduced evaluation"
437 );
438
439 let mut eval_point = sumcheck_output.challenges.clone();
441 eval_point.reverse();
442 assert_eq!(evaluate(&sparse_dense, &eval_point), multilinear_evals[0]);
443 assert_eq!(evaluate(dense, &eval_point), multilinear_evals[1]);
444 assert_eq!(output.challenges, sumcheck_output.challenges);
445 }
446
447 #[test]
448 fn a_factored_weight_proves_the_same_sum_as_the_table_it_stands_for() {
449 let mut rng = StdRng::seed_from_u64(7);
462
463 let factors = [2usize, 1, 2]
464 .iter()
465 .map(|&log_len| random_field_buffer::<P>(&mut rng, log_len))
466 .collect::<Vec<_>>();
467 let n_vars = 5;
468
469 let factored = FactoredMultilinear::new(factors);
471 let table = FieldBuffer::<P>::from_values(
472 &(0..1usize << n_vars)
473 .map(|index| factored.get(index))
474 .collect::<Vec<_>>(),
475 );
476
477 let sparse = random_sparse(&mut rng, n_vars, 12);
478 let sum = sparse
479 .iter()
480 .map(|&(index, value)| value * table.get(index))
481 .sum::<F>();
482 assert_ne!(sum, F::ZERO);
484
485 let transcripts = [
487 {
488 let prover = SparseDenseProductSumcheckProver::new(sparse.clone(), factored, sum);
489 let mut transcript = ProverTranscript::new(StdChallenger::default());
490 let output = prove_single(prover, &mut transcript);
491 (output.multilinear_evals, transcript.finalize())
492 },
493 {
494 let prover = SparseDenseProductSumcheckProver::new(
495 sparse.clone(),
496 FactoredMultilinear::new([table.clone()]),
497 sum,
498 );
499 let mut transcript = ProverTranscript::new(StdChallenger::default());
500 let output = prove_single(prover, &mut transcript);
501 (output.multilinear_evals, transcript.finalize())
502 },
503 ];
504
505 assert_eq!(transcripts[0].1, transcripts[1].1, "the two forms must prove the same rounds");
507 assert_eq!(transcripts[0].0, transcripts[1].0, "and reduce to the same evaluations");
508
509 prove_verify(sparse, &table);
511 }
512
513 #[test]
514 fn one_column_matches_the_single_column_prover() {
515 let mut rng = StdRng::seed_from_u64(17);
530
531 let n_vars = 5;
532 let weight = FactoredMultilinear::new([random_field_buffer::<P>(&mut rng, n_vars)]);
533 let sparse = random_sparse(&mut rng, n_vars, 14);
534 let sum = sparse
535 .iter()
536 .map(|&(index, _)| index)
537 .zip(sparse.iter().map(|&(_, value)| value))
538 .map(|(index, value)| value * weight.get(index))
539 .sum::<F>();
540 assert_ne!(sum, F::ZERO, "a vacuous claim would prove nothing");
541
542 let single = {
543 let prover = SparseDenseProductSumcheckProver::new(sparse.clone(), weight.clone(), sum);
544 let mut transcript = ProverTranscript::new(StdChallenger::default());
545 let output = batch_prove(vec![prover], &mut transcript);
546 (output.multilinear_evals[0].clone(), transcript.finalize())
547 };
548 let multi = {
549 let prover = SparseMultiDenseProductSumcheckProver::new(
550 sparse,
551 vec![weight],
552 std::slice::from_ref(&sum),
553 );
554 let mut transcript = ProverTranscript::new(StdChallenger::default());
555 let output = batch_prove(vec![prover], &mut transcript);
556 (output.multilinear_evals[0].clone(), transcript.finalize())
557 };
558
559 assert_eq!(single.1, multi.1, "the two provers must prove the same rounds");
560 assert_eq!(single.0, multi.0, "and reduce to the same evaluations");
561 }
562
563 #[test]
564 fn the_sparse_column_is_stored_once_however_many_dense_columns_ride_it() {
565 let mut rng = StdRng::seed_from_u64(19);
579
580 let n_vars = 4;
581 let sparse = random_sparse(&mut rng, n_vars, 10);
582 let weights = (0..3)
583 .map(|_| FactoredMultilinear::new([random_field_buffer::<P>(&mut rng, n_vars)]))
584 .collect::<Vec<_>>();
585 let sums = weights
586 .iter()
587 .map(|weight| {
588 sparse
589 .iter()
590 .map(|&(index, value)| value * weight.get(index))
591 .sum::<F>()
592 })
593 .collect::<Vec<_>>();
594
595 let prover = SparseMultiDenseProductSumcheckProver::new(sparse, weights, &sums);
596 let mut transcript = ProverTranscript::new(StdChallenger::default());
597 let output = batch_prove(vec![prover], &mut transcript);
598
599 let evals = &output.multilinear_evals[0];
601 assert_eq!(evals.len(), 4);
602
603 let sparse_eval = evals[0];
606 assert_ne!(sparse_eval, F::ZERO);
607 assert!(evals[1..].iter().all(|&dense| dense != sparse_eval));
608 }
609
610 #[test]
611 fn an_arena_backed_weight_proves_the_same_claim_as_a_heap_backed_one() {
612 let mut rng = StdRng::seed_from_u64(29);
623 let alloc = GlobalAllocator;
624
625 let n_vars = 5;
626 let scalars = (0..1usize << n_vars)
627 .map(|_| F::random(&mut rng))
628 .collect::<Vec<_>>();
629 let sparse = random_sparse(&mut rng, n_vars, 13);
630
631 let heap = FactoredMultilinear::<P>::new([FieldBuffer::<P>::from_values(&scalars)]);
632 let sum = sparse
633 .iter()
634 .map(|&(index, value)| value * heap.get(index))
635 .sum::<F>();
636 assert_ne!(sum, F::ZERO);
638
639 let prove_over = |weight| {
640 let prover = SparseDenseProductSumcheckProver::new(sparse.clone(), weight, sum);
641 let mut transcript = ProverTranscript::new(StdChallenger::default());
642 let output = prove_single(prover, &mut transcript);
643 (output.multilinear_evals, transcript.finalize())
644 };
645
646 let from_heap = prove_over(heap);
647 let from_arena = {
648 let weight =
649 FactoredMultilinear::new([FieldBuffer::<P>::from_values_in(&alloc, &scalars)]);
650 let prover = SparseDenseProductSumcheckProver::new(sparse, weight, sum);
651 let mut transcript = ProverTranscript::new(StdChallenger::default());
652 let output = prove_single(prover, &mut transcript);
653 (output.multilinear_evals, transcript.finalize())
654 };
655
656 assert_eq!(from_heap.1, from_arena.1, "both storages must prove the same rounds");
657 assert_eq!(from_heap.0, from_arena.0, "and reduce to the same evaluations");
658 }
659
660 fn random_sparse(
662 mut rng: impl rand::Rng,
663 n_vars: usize,
664 n_entries: usize,
665 ) -> Vec<SparseEntry<F>> {
666 (0..n_entries)
667 .map(|_| (rng.random_range(0..1 << n_vars), F::random(&mut rng)))
668 .collect()
669 }
670
671 proptest! {
672 #![proptest_config(ProptestConfig::with_cases(32))]
673
674 #[test]
681 fn prove_verify_matches_the_materialized_multilinears(
682 n_vars in 1usize..=8,
683 n_entries in 0usize..=100,
684 seed in any::<u64>(),
685 ) {
686 let mut rng = StdRng::seed_from_u64(seed);
687
688 let sparse = random_sparse(&mut rng, n_vars, n_entries);
689 let dense = random_field_buffer::<P>(&mut rng, n_vars);
690 prove_verify(sparse, &dense);
691 }
692 }
693
694 #[test]
697 fn test_sparse_dense_product_sumcheck_repeated_indices() {
698 let n_vars = 6;
699 let mut rng = StdRng::seed_from_u64(0);
700
701 let sparse = (0..10).map(|_| (23, F::random(&mut rng))).collect();
702 let dense = random_field_buffer::<P>(&mut rng, n_vars);
703 prove_verify(sparse, &dense);
704 }
705}