1use std::{iter, marker::PhantomData};
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, Divisible, PackedField};
9use binius_iop_prover::{
10 channel::IOPProverChannel,
11 logup_star::{self, Looker},
12};
13use binius_ip::prodcheck::MultilinearEvalClaim;
14use binius_ip_prover::{
15 channel::IPProverChannel,
16 prodcheck::{self, ProdcheckProver},
17 sumcheck::{
18 MleToSumCheckDecorator,
19 batch::{BatchSumcheckOutput, batch_prove, batch_prove_and_write_evals},
20 bivariate_product_mle,
21 multilinear_eval::multilinear_eval_prover,
22 quadratic_mlecheck_prover,
23 selector_mle::{Claim, SelectorMlecheckProver},
24 },
25};
26use binius_math::{
27 FieldSlice, FieldVec,
28 field_buffer::FieldBuffer,
29 inner_product::inner_product_buffers,
30 multilinear::{
31 eq::{eq_ind_partial_eval, eq_ind_partial_eval_scalars},
32 evaluate::{evaluate, evaluate_inplace},
33 },
34};
35use binius_utils::{checked_arithmetics::log2_ceil_usize, rayon::prelude::*};
36use binius_verifier::protocols::intmul::common::{
37 IntMulOutput, LIMB_BITS, LOG_N_LIMBS, N_LIMB_COLUMNS, N_LIMBS, Phase1Output, Phase2Output,
38 Phase3Output, Phase4Output, frobenius_twist, limb_column_twists, twist_limb_claim,
39};
40use either::Either;
41use itertools::izip;
42
43use super::{
44 error::Error,
45 witness::{Witness, limb_index, two_valued_field_buffer},
46};
47use crate::fold_word::{BitAxisFolder, WordAxisFolder};
48
49pub fn prove<A, F, P, Channel>(
69 columns: [&[Word]; 4],
70 channel: &mut Channel,
71 alloc: &A,
72) -> Result<IntMulOutput<F>, Error>
73where
74 A: Allocator,
75 F: BinaryField<Underlier: Divisible<u64>>,
76 P: PackedField<Scalar = F>,
77 Channel: IOPProverChannel<P, A>,
78{
79 let [a, b, lo, hi] = columns;
80
81 let witness = tracing::debug_span!("Build IntMul witness")
82 .in_scope(|| Witness::<_, P>::new(alloc, a, b, lo, hi))?;
83
84 let mut prover = IntMulProver::new(0, channel, alloc);
85 Ok(prover.prove(witness))
86}
87
88pub struct IntMulProver<'a, 'alloc, A: Allocator, P, Channel> {
91 _p_marker: PhantomData<P>,
92
93 switchover: usize,
94 channel: &'a mut Channel,
95 alloc: &'alloc A,
97}
98
99impl<'a, 'alloc, A: Allocator, P, Channel> IntMulProver<'a, 'alloc, A, P, Channel> {
100 pub const fn new(switchover: usize, channel: &'a mut Channel, alloc: &'alloc A) -> Self {
101 Self {
102 _p_marker: PhantomData,
103 switchover,
104 channel,
105 alloc,
106 }
107 }
108}
109
110impl<'alloc, A, F, P, Channel> IntMulProver<'_, 'alloc, A, P, Channel>
111where
112 A: Allocator,
113 F: BinaryField<Underlier: Divisible<u64>>,
114 P: PackedField<Scalar = F>,
115 Channel: IOPProverChannel<P, A>,
116{
117 pub fn prove(&mut self, witness: Witness<'_, 'alloc, A, P>) -> IntMulOutput<F> {
144 let Witness {
145 a_exponents,
146 a_prodcheck,
147 a_root,
148 b_exponents,
149 b_leaves,
150 b_prodcheck,
151 b_root,
152 c_lo_exponents,
153 c_lo_prodcheck,
154 c_lo_root,
155 c_hi_exponents,
156 c_hi_prodcheck,
157 c_hi_root,
158 tables,
159 } = witness;
160
161 let n_vars = b_root.log_len();
164
165 let initial_eval_point = self.channel.sample_many(n_vars);
166
167 let exp_eval = tracing::debug_span!("Evaluate exponent root")
169 .in_scope(|| evaluate_inplace(b_root, &initial_eval_point));
170
171 self.channel.send_one(exp_eval);
172
173 let Phase1Output {
175 eval_point: phase1_eval_point,
176 b_leaves_evals,
177 } = self.phase1(&initial_eval_point, b_prodcheck, b_leaves.as_view(), exp_eval);
178
179 let Phase2Output {
181 twisted_eval_points,
182 twisted_evals,
183 } = tracing::debug_span!("IntMul phase2 (Frobenius twist)")
184 .in_scope(|| frobenius_twist(Word::LOG_BITS, &phase1_eval_point, &b_leaves_evals));
185
186 let Phase3Output {
188 eval_point: phase3_eval_point,
189 r_ib,
190 b_recomb,
191 gpow_a_eval,
192 gpow_c_lo_eval,
193 gpow_c_hi_eval,
194 } = self.phase3(
195 &twisted_eval_points,
196 &twisted_evals,
197 a_root,
198 b_exponents,
199 [c_lo_root, c_hi_root],
200 &initial_eval_point,
201 exp_eval,
202 );
203
204 let phase_4_output = self.phase4(
206 &phase3_eval_point,
207 (gpow_a_eval, a_prodcheck),
208 (gpow_c_lo_eval, c_lo_prodcheck),
209 (gpow_c_hi_eval, c_hi_prodcheck),
210 [a_exponents, c_lo_exponents, c_hi_exponents],
211 &tables,
212 );
213
214 self.phase5(
216 &phase_4_output,
217 b_exponents,
218 &phase3_eval_point,
219 &r_ib,
220 b_recomb,
221 a_exponents,
222 c_lo_exponents,
223 c_hi_exponents,
224 tables[0].as_view(),
225 )
226 }
227
228 #[doc(hidden)] #[allow(clippy::too_many_arguments)]
230 pub fn phase5(
231 &mut self,
232 phase_4_output: &Phase4Output<F>,
233 b_exponents: &[Word],
234 b_eval_point: &[F],
235 r_ib: &[F],
236 b_recomb: F,
237 a_exponents: &[Word],
240 c_lo_exponents: &[Word],
241 c_hi_exponents: &[Word],
242 table: FieldSlice<'_, P>,
243 ) -> IntMulOutput<F> {
244 let alloc = self.alloc;
245 let n_vars = b_eval_point.len();
246 assert_eq!(phase_4_output.eval_point.len(), n_vars);
247
248 let twists = limb_column_twists();
252 let exponents = [a_exponents, c_lo_exponents, c_hi_exponents];
253 let limb_evals = [
254 &phase_4_output.a_limb_evals,
255 &phase_4_output.c_lo_limb_evals,
256 &phase_4_output.c_hi_limb_evals,
257 ];
258
259 let columns_guard = tracing::debug_span!("Gather lookup index columns").entered();
260 let index_columns = (0..N_LIMB_COLUMNS)
262 .into_par_iter()
263 .map(|j| {
264 let (tree, limb) = (j / N_LIMBS, j % N_LIMBS);
265 let n_padding = (1 << n_vars) - exponents[tree].len();
269 exponents[tree]
270 .iter()
271 .map(|&word| limb_index(word, limb))
272 .chain(iter::repeat_n(0, n_padding))
273 .collect::<Vec<_>>()
274 })
275 .collect::<Vec<_>>();
276 let twisted_claims = (0..N_LIMB_COLUMNS)
277 .map(|j| {
278 let (tree, limb) = (j / N_LIMBS, j % N_LIMBS);
279 twist_limb_claim(twists[j], &phase_4_output.eval_point, limb_evals[tree][limb])
280 })
281 .collect::<Vec<_>>();
282 drop(columns_guard);
283
284 let lookers = izip!(&index_columns, &twisted_claims)
289 .map(|(index, (twisted_point, twisted_eval))| Looker {
290 index,
291 eval_point: twisted_point,
292 eval_claim: *twisted_eval,
293 })
294 .collect::<Vec<_>>();
295 let log_cols = log2_ceil_usize(N_LIMB_COLUMNS);
296 let logup_proof = logup_star::prove_transparent(
300 [logup_star::TableLookup { table, lookers }],
301 self.channel,
302 self.alloc,
303 );
304 let [column_index_evals] = logup_proof.index_eval_claims.as_slice() else {
305 unreachable!("the reduction runs over the one power table")
306 };
307
308 let embed_guard = tracing::debug_span!("Build embedding table").entered();
311 let mut iota_table = Vec::with_capacity(1usize << LIMB_BITS);
312 iota_table.push(F::ZERO);
313 for row in 1..1usize << LIMB_BITS {
314 let low_bit_basis = F::basis(row.trailing_zeros() as usize);
315 iota_table.push(iota_table[row & (row - 1)] + low_bit_basis);
316 }
317 drop(embed_guard);
318
319 let index_content_point = logup_proof.index_eval_point.as_slice();
320
321 let rho = self.channel.sample_many(log_cols);
324 let mut padded_column_evals = column_index_evals.clone();
325 padded_column_evals.resize(1 << log_cols, F::ZERO);
326 let folded_index_claim =
327 evaluate(&FieldBuffer::<P>::from_values(&padded_column_evals), &rho);
328 let rho_tensor = eq_ind_partial_eval_scalars(&rho);
329 let fold_guard = tracing::debug_span!("Fold index columns by rho").entered();
330 let folded_column_scalars = (0..1usize << n_vars)
334 .into_par_iter()
335 .map(|i| {
336 izip!(&index_columns, &rho_tensor)
337 .map(|(column, &weight)| iota_table[column[i]] * weight)
338 .sum::<F>()
339 })
340 .collect::<Vec<_>>();
341 let folded_column = FieldBuffer::<P, _>::from_values_in(alloc, &folded_column_scalars);
342 drop(fold_guard);
343 let index_prover = MleToSumCheckDecorator::new(multilinear_eval_prover(
344 alloc,
345 folded_column,
346 index_content_point,
347 folded_index_claim,
348 ));
349
350 let binary_elements = [F::zero(), F::one()];
352
353 let a_0 = two_valued_field_buffer::<A, _, P>(alloc, 0, a_exponents, binary_elements);
355 let b_0 = two_valued_field_buffer::<A, _, P>(alloc, 0, b_exponents, binary_elements);
356 let c_lo_0 = two_valued_field_buffer::<A, _, P>(alloc, 0, c_lo_exponents, binary_elements);
357
358 let overflow_prover = MleToSumCheckDecorator::new(quadratic_mlecheck_prover(
361 alloc,
362 [a_0, b_0, c_lo_0],
363 |[a, b, c]| a * b - c,
364 |[a, b, _c]| a * b,
365 b_eval_point.to_vec(),
366 F::ZERO,
367 ));
368
369 assert!(b_exponents.len() <= 1 << n_vars);
373 let b_tensor = eq_ind_partial_eval_scalars(r_ib);
374 let b_folded = BitAxisFolder::new(&b_tensor).fold::<P, _>(alloc, b_exponents);
375 let b_sumcheck_prover = MleToSumCheckDecorator::new(multilinear_eval_prover(
376 alloc,
377 b_folded,
378 b_eval_point,
379 b_recomb,
380 ));
381
382 let batch_guard = tracing::debug_span!("Final batched sumcheck").entered();
383 let BatchSumcheckOutput {
384 mut challenges,
385 multilinear_evals: _,
386 } = batch_prove(
387 vec![
388 Either::Left(index_prover),
389 Either::Right(Either::Left(overflow_prover)),
390 Either::Right(Either::Right(b_sumcheck_prover)),
391 ],
392 self.channel,
393 );
394 drop(batch_guard);
395
396 challenges.reverse();
399 let r_out = challenges.as_slice();
400
401 let output_guard = tracing::debug_span!("Compute output bit evals").entered();
405 let folder = WordAxisFolder::<F>::new(r_out);
408 let [a_evals, b_evals, c_lo_evals, c_hi_evals] =
409 [a_exponents, b_exponents, c_lo_exponents, c_hi_exponents]
410 .into_par_iter()
411 .map(|exponents| folder.fold_par(exponents))
412 .collect::<Vec<_>>()
413 .try_into()
414 .expect("iterator over exact number of elements");
415 drop(output_guard);
416
417 self.channel.send_many(&a_evals);
418 self.channel.send_many(&c_lo_evals);
419 self.channel.send_many(&c_hi_evals);
420 self.channel.send_many(&b_evals);
421
422 IntMulOutput {
423 eval_point: r_out.to_vec(),
424 a_evals,
425 b_evals,
426 c_lo_evals,
427 c_hi_evals,
428 }
429 }
430}
431
432impl<'alloc, A, F, P, Channel> IntMulProver<'_, 'alloc, A, P, Channel>
433where
434 A: Allocator,
435 F: BinaryField,
436 P: PackedField<Scalar = F>,
437 Channel: IPProverChannel<F>,
438{
439 #[doc(hidden)] pub fn phase1(
441 &mut self,
442 eval_point: &[F],
443 b_prover: ProdcheckProver<'alloc, A, P>,
444 b_leaves: FieldSlice<'_, P>,
445 b_root_eval: F,
446 ) -> Phase1Output<F> {
447 let n_vars = eval_point.len();
448
449 let claim = MultilinearEvalClaim {
451 eval: b_root_eval,
452 point: eval_point.to_vec(),
453 };
454
455 let MultilinearEvalClaim {
457 eval: _,
458 point: reduced_point,
459 } = tracing::debug_span!("Variable-base product check")
460 .in_scope(|| b_prover.prove(claim, self.channel));
461
462 let (x_point, _z_suffix) = reduced_point.split_at(n_vars);
464
465 let leaf_guard = tracing::debug_span!("Compute base layer partial evals").entered();
467 let x_tensor = eq_ind_partial_eval(x_point);
468 let b_leaves_evals = b_leaves
469 .par_chunks(n_vars)
470 .map(|b_leaf| inner_product_buffers(&b_leaf, &x_tensor))
471 .collect::<Vec<_>>();
472 drop(leaf_guard);
473
474 self.channel.send_many(&b_leaves_evals);
476
477 Phase1Output {
478 eval_point: x_point.to_vec(),
479 b_leaves_evals,
480 }
481 }
482
483 #[doc(hidden)] #[allow(clippy::too_many_arguments)]
485 pub fn phase3(
486 &mut self,
487 twisted_eval_points: &[Vec<F>],
488 twisted_evals: &[F],
489 selector: FieldVec<P, A>,
490 b_exponents: &[Word],
491 c_lo_hi_roots: [FieldVec<P, A>; 2],
492 c_eval_point: &[F],
493 c_root_eval: F,
494 ) -> Phase3Output<F> {
495 let alloc = self.alloc;
496 let n_vars = selector.log_len();
497 assert!(
498 twisted_eval_points
499 .iter()
500 .all(|point| point.len() == n_vars)
501 );
502 assert!(b_exponents.len() <= 1 << n_vars);
503
504 let selector_claims = izip!(twisted_eval_points, twisted_evals)
505 .map(|(point, &value)| Claim {
506 point: point.clone(),
507 value,
508 })
509 .collect();
510
511 let gamma = self.channel.sample_many(Word::LOG_BITS);
518 let eq_weights = eq_ind_partial_eval_scalars(&gamma);
519 let mut b_bitmasks = bytemuck::cast_slice::<_, u64>(b_exponents).to_vec();
530 b_bitmasks.resize(1 << n_vars, 0);
531 let selector_prover = SelectorMlecheckProver::new(
532 selector,
533 selector_claims,
534 &b_bitmasks,
535 eq_weights,
536 self.switchover,
537 );
538
539 let c_root_sumcheck_prover =
540 bivariate_product_mle::new(alloc, c_lo_hi_roots, c_eval_point.to_vec(), c_root_eval);
541
542 let c_root_prover = MleToSumCheckDecorator::new(c_root_sumcheck_prover);
543
544 let provers = vec![Either::Left(selector_prover), Either::Right(c_root_prover)];
545 let sumcheck_guard = tracing::debug_span!("Batched selector + C-root sumcheck").entered();
546 let BatchSumcheckOutput {
547 mut challenges,
548 multilinear_evals,
549 } = batch_prove_and_write_evals(provers, self.channel);
550 challenges.reverse();
553 drop(sumcheck_guard);
554
555 let [mut selector_prover_evals, c_root_prover_evals] = multilinear_evals
556 .try_into()
557 .expect("batch_prove with two provers returns length-2 multilinear_evals");
558
559 assert_eq!(selector_prover_evals.len(), 1 + Word::BITS);
560
561 let gpow_a_eval = selector_prover_evals
562 .pop()
563 .expect("selector_prover_evals.len() > 0");
564 let b_evals = selector_prover_evals;
565 let [gpow_c_lo_eval, gpow_c_hi_eval] = c_root_prover_evals
566 .try_into()
567 .expect("c_root_prover with two multilinears returns two evals");
568
569 let r_ib = self.channel.sample_many(Word::LOG_BITS);
573 let b_recomb = evaluate(&FieldBuffer::<P>::from_values(&b_evals), &r_ib);
574
575 Phase3Output {
576 eval_point: challenges,
577 r_ib,
578 b_recomb,
579 gpow_a_eval,
580 gpow_c_lo_eval,
581 gpow_c_hi_eval,
582 }
583 }
584
585 #[doc(hidden)] #[allow(clippy::too_many_arguments)]
587 pub fn phase4(
588 &mut self,
589 eval_point: &[F],
590 (a_root_eval, a_prover): (F, ProdcheckProver<'alloc, A, P>),
591 (gpow_c_lo_eval, c_lo_prover): (F, ProdcheckProver<'alloc, A, P>),
592 (gpow_c_hi_eval, c_hi_prover): (F, ProdcheckProver<'alloc, A, P>),
593 exponents: [&[Word]; 3],
594 tables: &[FieldVec<P, A>],
595 ) -> Phase4Output<F> {
596 let n_vars = eval_point.len();
597
598 assert_eq!(a_prover.n_layers(), LOG_N_LIMBS);
600 assert_eq!(c_lo_prover.n_layers(), LOG_N_LIMBS);
601 assert_eq!(c_hi_prover.n_layers(), LOG_N_LIMBS);
602
603 let selector = self.channel.sample_many(log2_ceil_usize(3));
605
606 let prodcheck_guard = tracing::debug_span!("Batched constant-base prodcheck").entered();
610 let prodcheck::BatchProveOutput {
611 eval_point: reduced_point,
612 evals: _tree_evals,
613 } = prodcheck::batch_prove(
614 vec![a_prover, c_lo_prover, c_hi_prover],
615 vec![a_root_eval, gpow_c_lo_eval, gpow_c_hi_eval],
616 selector,
617 eval_point.to_vec(),
618 self.channel,
619 );
620 drop(prodcheck_guard);
621
622 let selector_len = log2_ceil_usize(3);
626 let (r_content, _r_limb) = reduced_point[selector_len..].split_at(n_vars);
627
628 let regather_guard = tracing::debug_span!("Regather per-limb evals").entered();
632 let twists = limb_column_twists();
633 let x_tensor = eq_ind_partial_eval(r_content);
634 let limb_evals = |tree: usize| {
635 (0..N_LIMBS)
636 .map(|limb| {
637 let table = &tables[twists[tree * N_LIMBS + limb] / LIMB_BITS];
638 let n_padding = (1 << n_vars) - exponents[tree].len();
641 let column_scalars = exponents[tree]
642 .iter()
643 .map(|&word| table.get(limb_index(word, limb)))
644 .chain(iter::repeat_n(table.get(0), n_padding))
645 .collect::<Vec<_>>();
646 let column = FieldBuffer::<P>::from_values(&column_scalars);
647 inner_product_buffers(&column, &x_tensor)
648 })
649 .collect::<Vec<_>>()
650 };
651 let a_limb_evals = limb_evals(0);
652 let c_lo_limb_evals = limb_evals(1);
653 let c_hi_limb_evals = limb_evals(2);
654 drop(regather_guard);
655 self.channel.send_many(&a_limb_evals);
656 self.channel.send_many(&c_lo_limb_evals);
657 self.channel.send_many(&c_hi_limb_evals);
658
659 Phase4Output {
660 eval_point: r_content.to_vec(),
661 a_limb_evals,
662 c_lo_limb_evals,
663 c_hi_limb_evals,
664 }
665 }
666}