Skip to main content

binius_prover/protocols/intmul/
witness.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use 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/// An integer multiplication protocol witness. Created from integer slices, consumed during
29/// proving.
30///
31/// The statement being proven is `a * b = c`, where `c` is represented as a pair `(c_lo, c_hi)`:
32/// `Word::BITS`-wide multiplicands with a double-wide product.
33///
34/// For each of `a`, `c_lo`, `c_hi` the exponent words split into `N_LIMBS` limbs, and the
35/// constant-base exponentiation factors as a product of `N_LIMBS` limb columns — column `l` reads
36/// the row `limb_l(e)` of the power table of the base $G^{2^{wl}}$ (where $w$ is the limb bit
37/// width). The columns are concatenated into one `(n_vars + LOG_N_LIMBS)`-variate buffer with the
38/// limb index in the high bits, and a [`ProdcheckProver`] is constructed over each:
39///  1) `a` and `c_lo` exponentiate the multiplicative group generator $G$
40///  2) `c_hi` exponentiates $G^{2^{2^m}}$
41///  3) `b` selects a variable base (the root of the `a` tree) per bit, over `Word::BITS` per-bit
42///     leaves as before
43///
44/// The shared power table $T\colon i \mapsto G^i$ over $2^w$ rows is retained for the Phase 5
45/// logup* lookup.
46///
47/// Protocol proves that ${(G^a)}^b = G^{c\\_lo} \times (G^{2^{2^m}})^{c\\_hi}$, which is equivalent
48/// to $a \times b = c$ modulo $2^{2^{m+1}} - 1$. The special case of `0 * 0 = 1` is handled
49/// separately.
50#[derive(Getters)]
51#[getset(get = "pub")]
52pub struct Witness<'a, 'alloc, A: Allocator, P: PackedField> {
53	/// The exponents for `a` (needed for the phase 5 lookup indices and parity zerocheck on
54	/// `a_0`).
55	#[getset(skip)]
56	pub a_exponents: &'a [Word],
57	/// Prodcheck prover for the `a` exponentiation tree (leaf layer retained).
58	pub a_prodcheck: ProdcheckProver<'alloc, A, P>,
59	/// The root of the `a` tree (product of all leaves element-wise); the `b` variable base.
60	pub a_root: FieldVec<P, A>,
61	/// The exponents for `b` (needed for phase 5).
62	#[getset(skip)]
63	pub b_exponents: &'a [Word],
64	/// Concatenated b leaves for prodcheck: [L_0, L_1, ..., L_{2^k-1}].
65	/// Has log_len = n_vars + Word::LOG_BITS.
66	pub b_leaves: FieldVec<P, A>,
67	/// The prover for the prodcheck reduction on b_leaves.
68	pub b_prodcheck: ProdcheckProver<'alloc, A, P>,
69	/// The root of the b tree (product of all leaves element-wise).
70	pub b_root: FieldVec<P, A>,
71	/// The exponents for `c_lo` (needed for the phase 5 lookup indices, parity zerocheck on
72	/// `c_lo_0`, and the raw per-bit output evaluations).
73	#[getset(skip)]
74	pub c_lo_exponents: &'a [Word],
75	/// Prodcheck prover for the `c_lo` exponentiation tree (leaf layer retained).
76	pub c_lo_prodcheck: ProdcheckProver<'alloc, A, P>,
77	/// The root of the `c_lo` tree.
78	pub c_lo_root: FieldVec<P, A>,
79	/// The exponents for `c_hi` (needed for the phase 5 lookup indices and raw per-bit output
80	/// evaluations).
81	#[getset(skip)]
82	pub c_hi_exponents: &'a [Word],
83	/// Prodcheck prover for the `c_hi` exponentiation tree (leaf layer retained).
84	pub c_hi_prodcheck: ProdcheckProver<'alloc, A, P>,
85	/// The root of the `c_hi` tree.
86	pub c_hi_root: FieldVec<P, A>,
87	/// The 2·N_LIMBS twisted power tables: `tables[s][i] = (G^{2^{ws}})^i`. Limb column `(t, l)`
88	/// is a gather from table `s(t, l)`; `tables[0]` is the shared table read by the Phase 5
89	/// logup* lookup.
90	pub tables: Vec<FieldVec<P, A>>,
91}
92
93// A manual `Clone` impl (rather than `#[derive(Clone)]`) so the bound lands on the pooled buffer
94// `A::Vec<P>` rather than on `A` and `P`. Holds for the `Vec`-backed `GlobalAllocator`.
95impl<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	/// Constructs a new integer multiplication witness from the statement.
126	///
127	/// The GKR prodcheck provers draw their layer buffers from `alloc`.
128	///
129	/// The four columns must have equal length, but need not have a power-of-two one: the reduction
130	/// runs over the constraint axis of `2^ceil(log2(n))` rows, and the rows past the columns' end
131	/// read as `Word::ZERO`. Every buffer built here still spans that whole axis, and each derives
132	/// its own value for a padding row rather than being handed one. In the product-check trees
133	/// that value is the multiplicative **identity**, not zero: a zero exponent word indexes row 0
134	/// of a power table, which holds `base^0 = 1`. So no tree's product moves.
135	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		// All statement slices should be of same length.
147		if [b, c_lo, c_hi]
148			.iter()
149			.any(|exponents| exponents.len() != a.len())
150		{
151			return Err(Error::ExponentLengthMismatch);
152		}
153
154		// The 2·N_LIMBS twisted power tables: tables[s][i] = (G^{2^{ws}})^i. The `a` and `c_lo`
155		// limb columns read tables 0..N_LIMBS; the `c_hi` limb columns (base
156		// G^{2^{Word::BITS}}) read tables N_LIMBS..2·N_LIMBS. Table 0 is the shared logup*
157		// table.
158		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		// Allocate the table buffers up front on this thread, so the parallel region only fills
165		// them — no allocator traffic inside the rayon closures.
166		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		// Build the per-limb leaf layers and their prodcheck provers. Each prover's products
177		// layer is the corresponding tree root.
178		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		// Compute b_leaves as concatenated leaves for prodcheck; the variable base is the `a` root.
191		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		// The prodcheck prover folds its leaf layer in place, and phase 1 reads the leaves again
195		// afterwards, so the prover gets a clone.
196		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
222/// Number of packed columns filled as one independent group before chaining to the next block.
223///
224/// The power chain `base^i` is inherently serial; striding the fill into blocks of
225/// `2^LOG_STRIDE` independent packed multiplies exposes instruction-level parallelism so the
226/// multiplier pipeline stays busy. `4` keeps a block (16 packed elements) comfortably in registers.
227const LOG_STRIDE: usize = 4;
228
229/// The first `count` powers of `base` (`base^0 .. base^(count-1)`) plus the next power
230/// `base^count`.
231fn 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
241/// Build the power table of `base` with `2^log_size` rows: row `i` holds `base^i`.
242pub 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
251/// Fill `buffer` with the power table of `base` and wrap it as a [`FieldBuffer`] of `2^log_size`
252/// rows: row `i` holds `base^i`.
253///
254/// The caller supplies the backing buffer, so its allocation can be hoisted out of a hot or
255/// parallel region; `buffer` is cleared and refilled here without reallocating.
256///
257/// # Preconditions
258///
259/// * `buffer.capacity()` must be at least `1 << log_size.saturating_sub(P::LOG_WIDTH)`.
260fn 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	// The first block spans `2^log_block` scalars = `2^LOG_STRIDE` packed elements.
274	let log_block = P::LOG_WIDTH + LOG_STRIDE;
275
276	// Small tables (at most one block) don't benefit from striding; build them sequentially. This
277	// also covers `log_size < P::LOG_WIDTH` (a single partially-filled packed element).
278	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	// Build the first block's `2^log_block` sequential powers, packed into `2^LOG_STRIDE` elements;
289	// `incr = base^(2^log_block)` is the per-block increment.
290	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	// Fill each remaining packed element from the one a block earlier; the `block_len` multiplies
299	// within a block are independent, so only the block-to-block step carries a dependency.
300	let block_len = 1 << LOG_STRIDE; // packed elements per block
301	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
309/// Extract limb `l` of an exponent word as a table row index.
310pub(super) const fn limb_index(word: Word, limb: usize) -> usize {
311	((word.0 >> (limb * LIMB_BITS)) & ((1 << LIMB_BITS) - 1)) as usize
312}
313
314/// Build the concatenated per-limb leaf columns for a constant-base GKR exponentiation tree.
315///
316/// Column `l` has entries `tables[l][limb_l(e)]` — the limb-`l` exponentiation of the
317/// corresponding word. The columns are concatenated into one `(n_vars + LOG_N_LIMBS)`-variate
318/// buffer with the limb index in the high bits, so the prodcheck's node reductions pair column `z`
319/// with column `z + N_LIMBS/2`.
320fn 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	// `exponents` need not fill the constraint axis. Its rows past the end are `Word::ZERO`, which
329	// indexes row 0 of every table — and that row holds `base^0 = 1`. So a padding row contributes
330	// the multiplicative identity to the product check, leaving each tree's product unchanged.
331	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/// Compute concatenated b_leaves for prodcheck.
348///
349/// Each leaf `L_z` contains: if bit z of `exponents[i]` is set then `bases[i]^{2^z}` else 1
350/// The leaves are concatenated: `[L_0, L_1, ..., L_{2^k-1}]`
351///
352/// The leaves are drawn from `alloc`: the prodcheck prover folds them in place, so they are a
353/// working buffer for the whole reduction.
354#[doc(hidden)] // exposed for benchmarking (`benches/intmul.rs`), not a stable API
355pub 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		// Parallel optimized path
369		return compute_b_leaves_parallel(alloc, bases, exponents);
370	}
371
372	// Fallback: bases is too small to parallelize (n_vars < P::LOG_WIDTH)
373	let mut out = FieldBuffer::zeros_in(alloc, n_vars + Word::LOG_BITS);
374	let n_elems = 1 << n_vars;
375
376	// Rows past the columns' end are `Word::ZERO`: every bit is clear, so each of their leaves is
377	// `F::ONE` — the identity the product check needs from a padding row.
378	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			// Branchless select of `base` when bit `z` is set, else `F::ONE`: on the selected lane
384			// `mask` is all-ones so `select` keeps `base - 1` and the `+ 1` restores `base`; on the
385			// unselected lane `select` yields `0` and the `+ 1` gives `F::ONE`.
386			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
396/// Parallel implementation of compute_b_leaves for when bases is large enough to parallelize.
397fn 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		// An allocator may hand back more capacity than requested, so take exactly the packed
416		// length the strided view expects.
417		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				// The columns' rows this packed element covers. Rows past their end are absent, and
428				// `make_mask` reads a lane it is not given as clear, so those lanes' leaves are all
429				// `F::ONE` — the identity the product check needs from a padding row.
430				let start = (packed_index * P::WIDTH).min(exponents.len());
431				let exp_chunk = &exponents[start..(start + P::WIDTH).min(exponents.len())];
432
433				// Keep base as packed element for efficient squaring
434				let mut packed_base = *packed_base;
435
436				for z in 0..height {
437					// Branchless select, per lane, of `base` when bit `z` of the lane's exponent is
438					// set, else `F::ONE`. On selected lanes `mask` is all-ones so `select` keeps
439					// `base - 1` and the `+ 1` restores `base`; on unselected lanes `select` yields
440					// `0` and the `+ 1` gives `F::ONE`.
441					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					// Square packed base for next iteration
445					packed_base = packed_base.square();
446				}
447			});
448	}
449
450	// SAFETY: All elements initialized in the parallel loop above
451	unsafe { out_vec.set_len(total) };
452
453	FieldBuffer::new(n_vars + Word::LOG_BITS, out_vec)
454}
455
456/// Compute the per-vertex bivariate product of two equally sized field buffers.
457pub 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
470/// Constructs a field buffer with values selected from `elements` based on the bit values
471/// of `exponents`.
472pub 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	// Select `elements[1]` if bit `bit_offset` of the word is set, else `elements[0]`. A row past
487	// the columns' end is `Word::ZERO`, whose bits are all clear, so it selects `elements[0]`.
488	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	// The packed elements `exponents` fills whole. Rounding down to a multiple of `P::WIDTH` keeps
494	// every lane of this loop in range, so it packs without a per-lane bounds check.
495	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	// The columns' trailing words share a packed element with the start of the padding.
503	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	// The rest of the constraint axis is padding.
514	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		// The variable-base `b`-exponent tree root must equal the full product `c` root
530		// (`c_lo_root * c_hi_root`); this equality is what lets the prover reuse `b_root` in place
531		// of a separately stored `c_root`.
532		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)]; // 2*3 = 6
541		let c_hi = [Word::from_u64(0)]; // no high bits
542
543		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)]; // 2^32 * 2^33 = 2^65, which is 2 in the high 64 bits
554
555		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	/// Directly checks `compute_b_leaves` against its specification — leaf `z` holds
591	/// `bases[i]^(2^z)` where bit `z` of `exponents[i]` is set, else `F::ONE` — over both the
592	/// parallel path (`n_vars >= P::LOG_WIDTH`) and the scalar fallback (`n_vars < P::LOG_WIDTH`).
593	#[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		// `Packed128b` has `LOG_WIDTH == 2`: `n_vars = 0` exercises the scalar fallback and
602		// `n_vars = 4` the parallel path.
603		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	/// Checks `power_table` against a sequential reference (`row i == base^i`) across both the
635	/// small sequential fallback (`log_size <= P::LOG_WIDTH + LOG_STRIDE`) and the strided packed
636	/// path.
637	#[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		// `Packed128b` has `LOG_WIDTH == 2`, so `log_block == 6`: sizes up to 6 take the sequential
647		// fallback; 7 and 10 (multiple blocks) exercise the strided path.
648		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}