1use std::{iter, mem::MaybeUninit, ops::Deref};
5
6use binius_compute::{Allocator, VecLike};
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, PackedField};
9use binius_ip_prover::prodcheck::ProdcheckProver;
10use binius_math::{
11 FieldVec,
12 field_buffer::{FieldBuffer, FieldSlice},
13};
14use binius_utils::{
15 checked_arithmetics::log2_ceil_usize,
16 rayon::{
17 prelude::*,
18 task_size::{IndexedParallelIteratorExt, WorkPerItem},
19 },
20 strided_array::StridedArray2DViewMut,
21};
22use binius_verifier::protocols::intmul::common::{LIMB_BITS, LOG_N_LIMBS, N_LIMBS};
23use getset::Getters;
24use itertools::iterate;
25
26use super::error::Error;
27
28#[derive(Getters)]
51#[getset(get = "pub")]
52pub struct Witness<'a, 'alloc, A: Allocator, P: PackedField> {
53 #[getset(skip)]
56 pub a_exponents: &'a [Word],
57 pub a_prodcheck: ProdcheckProver<'alloc, A, P>,
59 pub a_root: FieldVec<P, A>,
61 #[getset(skip)]
63 pub b_exponents: &'a [Word],
64 pub b_leaves: FieldVec<P, A>,
67 pub b_prodcheck: ProdcheckProver<'alloc, A, P>,
69 pub b_root: FieldVec<P, A>,
71 #[getset(skip)]
74 pub c_lo_exponents: &'a [Word],
75 pub c_lo_prodcheck: ProdcheckProver<'alloc, A, P>,
77 pub c_lo_root: FieldVec<P, A>,
79 #[getset(skip)]
82 pub c_hi_exponents: &'a [Word],
83 pub c_hi_prodcheck: ProdcheckProver<'alloc, A, P>,
85 pub c_hi_root: FieldVec<P, A>,
87 pub tables: Vec<FieldVec<P, A>>,
91}
92
93impl<A: Allocator, P: PackedField> Clone for Witness<'_, '_, A, P>
96where
97 A::Vec<P>: Clone,
98{
99 fn clone(&self) -> Self {
100 Self {
101 a_exponents: self.a_exponents,
102 a_prodcheck: self.a_prodcheck.clone(),
103 a_root: self.a_root.clone(),
104 b_exponents: self.b_exponents,
105 b_leaves: self.b_leaves.clone(),
106 b_prodcheck: self.b_prodcheck.clone(),
107 b_root: self.b_root.clone(),
108 c_lo_exponents: self.c_lo_exponents,
109 c_lo_prodcheck: self.c_lo_prodcheck.clone(),
110 c_lo_root: self.c_lo_root.clone(),
111 c_hi_exponents: self.c_hi_exponents,
112 c_hi_prodcheck: self.c_hi_prodcheck.clone(),
113 c_hi_root: self.c_hi_root.clone(),
114 tables: self.tables.clone(),
115 }
116 }
117}
118
119impl<'a, 'alloc, A, F, P> Witness<'a, 'alloc, A, P>
120where
121 A: Allocator,
122 F: BinaryField,
123 P: PackedField<Scalar = F>,
124{
125 pub fn new(
136 alloc: &'alloc A,
137 a: &'a [Word],
138 b: &'a [Word],
139 c_lo: &'a [Word],
140 c_hi: &'a [Word],
141 ) -> Result<Self, Error> {
142 const {
143 assert!(2 * Word::BITS <= F::N_BITS, "F must be wide enough to hold a 128-bit product");
144 }
145
146 if [b, c_lo, c_hi]
148 .iter()
149 .any(|exponents| exponents.len() != a.len())
150 {
151 return Err(Error::ExponentLengthMismatch);
152 }
153
154 let power_table_scope = tracing::debug_span!("Build power tables").entered();
159 let g = F::MULTIPLICATIVE_GENERATOR;
160 let bases = iterate(g, |g| g.square())
161 .step_by(LIMB_BITS)
162 .take(2 * N_LIMBS)
163 .collect::<Vec<_>>();
164 let packed_len = 1 << LIMB_BITS.saturating_sub(P::LOG_WIDTH);
167 let buffers = iter::repeat_with(|| alloc.alloc::<P>(packed_len))
168 .take(bases.len())
169 .collect::<Vec<_>>();
170 let tables = (bases, buffers)
171 .into_par_iter()
172 .map(|(base, buffer)| power_table_into(LIMB_BITS, base, buffer))
173 .collect::<Vec<_>>();
174 drop(power_table_scope);
175
176 let fixed_base_tree_scope =
179 tracing::debug_span!("Compute fixed-base prodcheck layers").entered();
180 let a_leaves = limb_leaves::<A, F, P>(alloc, &tables[..N_LIMBS], a);
181 let (a_prodcheck, a_root) = ProdcheckProver::new(LOG_N_LIMBS, alloc, a_leaves);
182
183 let c_lo_leaves = limb_leaves::<A, F, P>(alloc, &tables[..N_LIMBS], c_lo);
184 let (c_lo_prodcheck, c_lo_root) = ProdcheckProver::new(LOG_N_LIMBS, alloc, c_lo_leaves);
185
186 let c_hi_leaves = limb_leaves::<A, F, P>(alloc, &tables[N_LIMBS..], c_hi);
187 let (c_hi_prodcheck, c_hi_root) = ProdcheckProver::new(LOG_N_LIMBS, alloc, c_hi_leaves);
188 drop(fixed_base_tree_scope);
189
190 let variable_base_tree_scope =
192 tracing::debug_span!("Compute variable-base prodcheck layers").entered();
193 let b_leaves = compute_b_leaves(alloc, a_root.as_view(), b);
194 let (b_prodcheck, b_root) = ProdcheckProver::new(
197 Word::LOG_BITS,
198 alloc,
199 FieldBuffer::from_view_in(alloc, b_leaves.as_view()),
200 );
201 drop(variable_base_tree_scope);
202
203 Ok(Self {
204 a_exponents: a,
205 a_prodcheck,
206 a_root,
207 b_exponents: b,
208 b_leaves,
209 b_prodcheck,
210 b_root,
211 c_lo_exponents: c_lo,
212 c_lo_prodcheck,
213 c_lo_root,
214 c_hi_exponents: c_hi,
215 c_hi_prodcheck,
216 c_hi_root,
217 tables,
218 })
219 }
220}
221
222const LOG_STRIDE: usize = 4;
228
229fn sequential_powers<F: Field>(count: usize, base: F) -> (Vec<F>, F) {
232 let mut scalars = Vec::with_capacity(count);
233 let mut row = F::ONE;
234 for _ in 0..count {
235 scalars.push(row);
236 row *= base;
237 }
238 (scalars, row)
239}
240
241pub fn power_table<F, P>(log_size: usize, base: F) -> FieldBuffer<P>
243where
244 F: Field,
245 P: PackedField<Scalar = F>,
246{
247 let buffer = Vec::with_capacity(1 << log_size.saturating_sub(P::LOG_WIDTH));
248 power_table_into(log_size, base, buffer)
249}
250
251fn power_table_into<F, P, V>(log_size: usize, base: F, mut buffer: V) -> FieldBuffer<P, V>
261where
262 F: Field,
263 P: PackedField<Scalar = F>,
264 V: VecLike<P>,
265{
266 let packed_len = 1 << log_size.saturating_sub(P::LOG_WIDTH);
267 assert!(
268 buffer.capacity() >= packed_len,
269 "precondition: buffer capacity must cover the packed table length"
270 );
271 buffer.clear();
272
273 let log_block = P::LOG_WIDTH + LOG_STRIDE;
275
276 if log_size <= log_block {
279 let (scalars, _) = sequential_powers(1 << log_size, base);
280 buffer.extend(
281 scalars
282 .chunks(P::WIDTH)
283 .map(|chunk| P::from_scalars(chunk.iter().copied())),
284 );
285 return FieldBuffer::new(log_size, buffer);
286 }
287
288 let (block_scalars, incr_scalar) = sequential_powers(1 << log_block, base);
291 let incr = P::broadcast(incr_scalar);
292 buffer.extend(
293 block_scalars
294 .chunks(P::WIDTH)
295 .map(|chunk| P::from_scalars(chunk.iter().copied())),
296 );
297
298 let block_len = 1 << LOG_STRIDE; for i in block_len..packed_len {
302 let next = buffer[i - block_len] * incr;
303 buffer.push(next);
304 }
305
306 FieldBuffer::new(log_size, buffer)
307}
308
309pub(super) const fn limb_index(word: Word, limb: usize) -> usize {
311 ((word.0 >> (limb * LIMB_BITS)) & ((1 << LIMB_BITS) - 1)) as usize
312}
313
314fn limb_leaves<A, F, P>(alloc: &A, tables: &[FieldVec<P, A>], exponents: &[Word]) -> FieldVec<P, A>
321where
322 A: Allocator,
323 F: Field,
324 P: PackedField<Scalar = F>,
325{
326 assert_eq!(tables.len(), N_LIMBS);
327
328 let n_vars = log2_ceil_usize(exponents.len());
332 let n_padding = (1 << n_vars) - exponents.len();
333 let scalars = (0..N_LIMBS)
334 .flat_map(|limb| {
335 let table = &tables[limb];
336 exponents
337 .iter()
338 .map(move |&word| table.get(limb_index(word, limb)))
339 .chain(iter::repeat_n(table.get(0), n_padding))
340 })
341 .collect::<Vec<_>>();
342
343 debug_assert_eq!(scalars.len(), 1 << (n_vars + LOG_N_LIMBS));
344 FieldBuffer::from_values_in(alloc, &scalars)
345}
346
347#[doc(hidden)] pub fn compute_b_leaves<A, F, P>(
356 alloc: &A,
357 bases: FieldSlice<'_, P>,
358 exponents: &[Word],
359) -> FieldVec<P, A>
360where
361 A: Allocator,
362 F: Field,
363 P: PackedField<Scalar = F>,
364{
365 let n_vars = bases.log_len();
366
367 if P::LOG_WIDTH <= n_vars {
368 return compute_b_leaves_parallel(alloc, bases, exponents);
370 }
371
372 let mut out = FieldBuffer::zeros_in(alloc, n_vars + Word::LOG_BITS);
374 let n_elems = 1 << n_vars;
375
376 let padding = iter::repeat_n(Word::ZERO, n_elems - exponents.len());
379 let exponents = exponents.iter().copied().chain(padding);
380
381 for (i, (mut base, exp)) in iter::zip(bases.iter_scalars(), exponents).enumerate() {
382 for z in 0..Word::BITS {
383 let mask = F::make_mask(iter::once(exp.extract_bit(z)));
387 out.set(z * n_elems + i, F::ONE + (base - F::ONE).select(&mask));
388
389 base = base.square();
390 }
391 }
392
393 out
394}
395
396fn compute_b_leaves_parallel<A, F, P>(
398 alloc: &A,
399 bases: FieldSlice<'_, P>,
400 exponents: &[Word],
401) -> FieldVec<P, A>
402where
403 A: Allocator,
404 F: Field,
405 P: PackedField<Scalar = F>,
406{
407 let n_vars = bases.log_len();
408 let n_packed = bases.as_ref().len();
409 let height = Word::BITS;
410 let total = n_packed * height;
411
412 let mut out_vec = alloc.alloc::<P>(total);
413
414 {
415 let spare: &mut [MaybeUninit<P>] = &mut out_vec.spare_capacity_mut()[..total];
418
419 let mut strided = StridedArray2DViewMut::without_stride(spare, height, n_packed)
420 .expect("dimensions match capacity");
421
422 let ones = P::broadcast(F::ONE);
423 (strided.par_iter_cols(), bases.as_ref())
424 .into_par_iter()
425 .enumerate()
426 .for_each(|(packed_index, (mut col, packed_base))| {
427 let start = (packed_index * P::WIDTH).min(exponents.len());
431 let exp_chunk = &exponents[start..(start + P::WIDTH).min(exponents.len())];
432
433 let mut packed_base = *packed_base;
435
436 for z in 0..height {
437 let mask = P::make_mask(exp_chunk.iter().map(|&exp| exp.extract_bit(z)));
442 col[z].write(ones + (packed_base - ones).select(&mask));
443
444 packed_base = packed_base.square();
446 }
447 });
448 }
449
450 unsafe { out_vec.set_len(total) };
452
453 FieldBuffer::new(n_vars + Word::LOG_BITS, out_vec)
454}
455
456pub fn buffer_bivariate_product<P: PackedField, Data: Deref<Target = [P]>>(
458 a: &FieldBuffer<P, Data>,
459 b: &FieldBuffer<P, Data>,
460) -> FieldBuffer<P> {
461 assert_eq!(a.len(), b.len());
462 let product = (a.as_ref(), b.as_ref())
463 .into_par_iter()
464 .with_min_task(WorkPerItem::FieldMuls)
465 .map(|(&a, &b)| a * b)
466 .collect::<Vec<P>>();
467 FieldBuffer::new(a.log_len(), product)
468}
469
470pub fn two_valued_field_buffer<A, F, P>(
473 alloc: &A,
474 bit_offset: usize,
475 exponents: &[Word],
476 elements: [F; 2],
477) -> FieldVec<P, A>
478where
479 A: Allocator,
480 F: Field,
481 P: PackedField<Scalar = F>,
482{
483 let n_vars = log2_ceil_usize(exponents.len());
484 let packed_len = 1 << n_vars.saturating_sub(P::LOG_WIDTH);
485
486 let select = |&word: &Word| elements[word.extract_bit(bit_offset) as usize];
489 let padding = elements[0];
490
491 let mut values = alloc.alloc::<P>(packed_len);
492
493 let mut chunks = exponents.chunks_exact(P::WIDTH);
496 values.extend(
497 chunks
498 .by_ref()
499 .map(|chunk| P::from_scalars(chunk.iter().map(select))),
500 );
501
502 let tail = chunks.remainder();
504 if !tail.is_empty() {
505 values.push(P::from_scalars(
506 tail.iter()
507 .map(select)
508 .chain(iter::repeat(padding))
509 .take(P::WIDTH),
510 ));
511 }
512
513 values.resize(packed_len, P::broadcast(padding));
515
516 FieldBuffer::new(n_vars, values)
517}
518
519#[cfg(test)]
520mod tests {
521 use binius_compute::GlobalAllocator;
522 use binius_math::test_utils::Packed128b;
523
524 use super::*;
525
526 type P = Packed128b;
527
528 fn check_consistency<A: Allocator, P: PackedField>(witness: &Witness<'_, '_, A, P>) {
529 let c_root = buffer_bivariate_product(witness.c_lo_root(), witness.c_hi_root());
533 assert_eq!(witness.b_root().as_view(), c_root.as_view());
534 }
535
536 #[test]
537 fn test_forwards() {
538 let a = [Word::from_u64(2)];
539 let b = [Word::from_u64(3)];
540 let c_lo = [Word::from_u64(6)]; let c_hi = [Word::from_u64(0)]; let alloc = GlobalAllocator;
544 let witness = Witness::<_, P>::new(&alloc, &a, &b, &c_lo, &c_hi).unwrap();
545 check_consistency(&witness);
546 }
547
548 #[test]
549 fn test_forwards_larger() {
550 let a = [Word::from_u64(1 << 32)];
551 let b = [Word::from_u64(1 << 33)];
552 let c_lo = [Word::from_u64(0)];
553 let c_hi = [Word::from_u64(2)]; let alloc = GlobalAllocator;
556 let witness = Witness::<_, P>::new(&alloc, &a, &b, &c_lo, &c_hi).unwrap();
557 check_consistency(&witness);
558 }
559
560 #[test]
561 fn test_forwards_multiple_random() {
562 use rand::prelude::*;
563
564 let mut rng = StdRng::seed_from_u64(0);
565
566 const VECTOR_SIZE: usize = 8;
567 let mut a = Vec::with_capacity(VECTOR_SIZE);
568 let mut b = Vec::with_capacity(VECTOR_SIZE);
569 let mut c_lo = Vec::with_capacity(VECTOR_SIZE);
570 let mut c_hi = Vec::with_capacity(VECTOR_SIZE);
571
572 for _ in 0..VECTOR_SIZE {
573 let a_i = rng.random_range(1..u64::MAX);
574 let b_i = rng.random_range(1..u64::MAX);
575
576 let full_result = (a_i as u128) * (b_i as u128);
577 let c_lo_i = full_result as u64;
578 let c_hi_i = (full_result >> 64) as u64;
579
580 a.push(Word::from_u64(a_i));
581 b.push(Word::from_u64(b_i));
582 c_lo.push(Word::from_u64(c_lo_i));
583 c_hi.push(Word::from_u64(c_hi_i));
584 }
585
586 let alloc = GlobalAllocator;
587 let witness = Witness::<_, P>::new(&alloc, &a, &b, &c_lo, &c_hi).unwrap();
588 check_consistency(&witness);
589 }
590 #[test]
594 fn compute_b_leaves_matches_spec() {
595 use binius_field::{Ghash128b, Random, arithmetic_traits::Square};
596 use rand::prelude::*;
597
598 type F = Ghash128b;
599
600 let mut rng = StdRng::seed_from_u64(1);
601 for n_vars in [0usize, 4] {
604 let n_elems = 1 << n_vars;
605 let base_scalars = (0..n_elems)
606 .map(|_| F::random(&mut rng))
607 .collect::<Vec<_>>();
608 let bases = FieldBuffer::<P>::from_values(&base_scalars);
609 let exponents = (0..n_elems)
610 .map(|_| Word::from_u64(rng.random()))
611 .collect::<Vec<_>>();
612
613 let leaves = compute_b_leaves::<_, F, P>(&GlobalAllocator, bases.as_view(), &exponents);
614
615 for (i, &base0) in base_scalars.iter().enumerate() {
616 let mut base = base0;
617 for z in 0..Word::BITS {
618 let expected = if exponents[i].extract_bit(z) {
619 base
620 } else {
621 F::ONE
622 };
623 assert_eq!(
624 leaves.get(z * n_elems + i),
625 expected,
626 "mismatch at n_vars={n_vars}, i={i}, z={z}"
627 );
628 base = base.square();
629 }
630 }
631 }
632 }
633
634 #[test]
638 fn power_table_matches_sequential() {
639 use binius_field::{Ghash128b, Random};
640 use rand::prelude::*;
641
642 type F = Ghash128b;
643
644 let mut rng = StdRng::seed_from_u64(2);
645 let base = F::random(&mut rng);
646 for log_size in [0usize, 3, 6, 7, 10] {
649 let table = power_table::<F, P>(log_size, base);
650 let mut expected = F::ONE;
651 for i in 0..1usize << log_size {
652 assert_eq!(table.get(i), expected, "mismatch at log_size={log_size}, i={i}");
653 expected *= base;
654 }
655 }
656 }
657}