Skip to main content

binius_prover/protocols/
binmul.rs

1// Copyright 2026 The Binius Developers
2
3use std::iter;
4
5use binius_compute::{Allocator, VecLike};
6use binius_core::word::Word;
7use binius_field::{BinaryField, PackedField};
8use binius_ip_prover::{
9	channel::IPProverChannel,
10	sumcheck::{ProveSingleOutput, prove_single_mlecheck, quadratic_mlecheck_prover},
11};
12use binius_math::{FieldVec, field_buffer::FieldBuffer};
13use binius_utils::{checked_arithmetics::log2_ceil_usize, rayon::prelude::*};
14pub use binius_verifier::protocols::binmul::BinMulOutput;
15
16use crate::fold_word::WordAxisFolder;
17
18/// Build a packed GHASH-field multilinear table from a `(lo, hi)` pair of word columns.
19///
20/// The scalar at hypercube index `i` is the field element carried by the pair `(lo[i], hi[i])`:
21/// `lo[i]` supplies the low 64 bits and `hi[i]` the high 64 bits of the 128-bit value.
22///
23/// The columns need not have a power-of-two length. The table spans the whole `2^n_vars` hypercube
24/// regardless, and an index at or past the columns' end reads as the zero field element. The
25/// padding therefore lives in this buffer, which is allocated at the hypercube's size either way,
26/// rather than in a `Vec<Word>` the caller has to materialize.
27///
28/// # Preconditions
29///
30/// * The two columns have equal length.
31fn build_table<A, F, P>(alloc: &A, lo: &[Word], hi: &[Word]) -> FieldVec<P, A>
32where
33	A: Allocator,
34	F: BinaryField + From<u128>,
35	P: PackedField<Scalar = F>,
36{
37	assert_eq!(lo.len(), hi.len());
38
39	let n_vars = log2_ceil_usize(lo.len());
40	let packed_len = 1 << n_vars.saturating_sub(P::LOG_WIDTH);
41	let elem =
42		|(&lo, &hi): (&Word, &Word)| F::from((lo.as_u64() as u128) | ((hi.as_u64() as u128) << 64));
43
44	let mut values = alloc.alloc::<P>(packed_len);
45
46	// The packed elements the columns fill whole. `chunks_exact` stops at the last of them, so
47	// every lane this loop packs is a real row.
48	let (lo_chunks, hi_chunks) = (lo.chunks_exact(P::WIDTH), hi.chunks_exact(P::WIDTH));
49	let (lo_tail, hi_tail) = (lo_chunks.remainder(), hi_chunks.remainder());
50	values.extend(
51		iter::zip(lo_chunks, hi_chunks)
52			.map(|(lo_chunk, hi_chunk)| P::from_scalars(iter::zip(lo_chunk, hi_chunk).map(elem))),
53	);
54
55	// The columns' trailing words share one packed element with the start of the padding. The lanes
56	// `from_scalars` is not given default to zero, which is exactly that padding.
57	if !lo_tail.is_empty() {
58		values.push(P::from_scalars(iter::zip(lo_tail, hi_tail).map(elem)));
59	}
60
61	// The rest of the hypercube is zero padding. A zero row satisfies `0 * 0 = 0`, so it
62	// contributes nothing to the zerocheck.
63	values.resize(packed_len, P::zero());
64
65	FieldBuffer::new(n_vars, values)
66}
67
68/// Prove the binary-field multiplication check (BinMul) reduction.
69///
70/// Proves $\widetilde{A}(x) \cdot \widetilde{B}(x) = \widetilde{C}(x)$ for every $x$ on the boolean
71/// hypercube $\mathbb{B}_\ell$ over the GHASH field, where each element is carried by a `(lo, hi)`
72/// pair of 64-bit words. See [`binius_verifier::protocols::binmul::verify`] for the protocol
73/// description and output shape.
74///
75/// The six `columns` are the `(lo, hi)` word pairs of the two multiplicands and the product, in the
76/// order `[a_lo, a_hi, b_lo, b_hi, c_lo, c_hi]`, all of equal length. That length need not be a
77/// power of two: the hypercube is $\mathbb{B}_\ell$ for $\ell = \lceil \log_2 n \rceil$, and a row
78/// at or past the columns' end reads as the zero field element. A zero row satisfies
79/// $0 \cdot 0 = 0$, so it contributes nothing to the zerocheck. The GHASH-field element for row $x$
80/// is
81/// $\langle\langle z_{\textsf{lo}}, z_{\textsf{hi}} \rangle\rangle = \sum_{i=0}^{63}
82/// z_{\textsf{lo},x,i} \cdot X^i + \sum_{i=0}^{63} z_{\textsf{hi},x,i} \cdot X^{64+i}$ for each of
83/// $z \in \{a, b, c\}$.
84pub fn prove<A, F, P, Channel>(
85	columns: [&[Word]; 6],
86	channel: &mut Channel,
87	alloc: &A,
88) -> BinMulOutput<F>
89where
90	A: Allocator,
91	F: BinaryField + From<u128>,
92	P: PackedField<Scalar = F>,
93	Channel: IPProverChannel<F>,
94{
95	let [a_lo, a_hi, b_lo, b_hi, c_lo, c_hi] = columns;
96
97	let n_vars = log2_ceil_usize(a_lo.len());
98	for column in [a_hi, b_lo, b_hi, c_lo, c_hi] {
99		assert_eq!(column.len(), a_lo.len());
100	}
101
102	// Build the packed GHASH-field multilinear tables A, B, C from the (lo, hi) word pairs.
103	let a = build_table::<A, F, P>(alloc, a_lo, a_hi);
104	let b = build_table::<A, F, P>(alloc, b_lo, b_hi);
105	let c = build_table::<A, F, P>(alloc, c_lo, c_hi);
106
107	// Sample the zerocheck challenge r_z.
108	let r_z = channel.sample_many(n_vars);
109
110	// Product zerocheck: 0 = sum_x eq(r_z, x) * (A(x) * B(x) - C(x)). The composition A * B - C has
111	// degree 2; the eq factor is folded internally by the MLE-check.
112	let prover = quadratic_mlecheck_prover(
113		alloc,
114		[a, b, c],
115		|[a, b, c]| a * b - c,
116		|[a, b, _c]| a * b,
117		r_z,
118		F::ZERO,
119	);
120
121	let ProveSingleOutput {
122		multilinear_evals: _,
123		challenges: mut eval_point,
124	} = prove_single_mlecheck(prover, channel);
125	// `prove_single_mlecheck` folds high-to-low, so reverse to obtain the shared output point r_x.
126	eval_point.reverse();
127
128	// Send the raw per-bit output evaluations of each word column at r_x. One folder is built for
129	// the shared point and reused across all six columns.
130	//
131	// Six columns cannot fill a machine on their own, so each column divides its own chunk axis
132	// rather than the six being run against each other.
133	let folder = WordAxisFolder::<F>::new(&eval_point);
134	let [
135		a_lo_evals,
136		a_hi_evals,
137		b_lo_evals,
138		b_hi_evals,
139		c_lo_evals,
140		c_hi_evals,
141	] = [a_lo, a_hi, b_lo, b_hi, c_lo, c_hi]
142		.into_par_iter()
143		.map(|exponents| folder.fold_par(exponents))
144		.collect::<Vec<_>>()
145		.try_into()
146		.expect("iterator over exact number of elements");
147
148	channel.send_many(&a_lo_evals);
149	channel.send_many(&a_hi_evals);
150	channel.send_many(&b_lo_evals);
151	channel.send_many(&b_hi_evals);
152	channel.send_many(&c_lo_evals);
153	channel.send_many(&c_hi_evals);
154
155	BinMulOutput {
156		eval_point,
157		a_lo_evals,
158		a_hi_evals,
159		b_lo_evals,
160		b_hi_evals,
161		c_lo_evals,
162		c_hi_evals,
163	}
164}
165
166#[cfg(test)]
167mod tests {
168	use binius_compute::GlobalAllocator;
169	use binius_core::word::Word;
170	use binius_field::{Ghash128b, PackedGhash2x128b, Random};
171	use binius_iop::channel::{OracleSpec, naive::NaiveVerifierChannel};
172	use binius_iop_prover::channel::naive::NaiveProverChannel;
173	use binius_math::{inner_product::inner_product_buffers, multilinear::eq::eq_ind_partial_eval};
174	use binius_transcript::{ProverTranscript, VerifierTranscript};
175	use binius_verifier::{
176		config::StdChallenger,
177		protocols::binmul::{BinMulOutput, verify},
178	};
179	use itertools::izip;
180	use rand::prelude::*;
181
182	use super::{log2_ceil_usize, prove};
183
184	type F = Ghash128b;
185	type P = PackedGhash2x128b;
186
187	/// The six word columns `(a_lo, a_hi, b_lo, b_hi, c_lo, c_hi)` of a BinMul witness.
188	type WordColumns = (Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>);
189
190	/// Evaluate the multilinear extension of a per-bit word column at a point, independently of the
191	/// prover. The point's `Word::LOG_BITS`-coordinate prefix selects the bit within a word; the
192	/// suffix selects the word (constraint) index.
193	fn evaluate_witness(words: &[Word], eval_point: &[F]) -> F {
194		let (prefix, suffix) = eval_point.split_at(Word::LOG_BITS);
195		let prefix_tensor = eq_ind_partial_eval::<F>(prefix);
196		let suffix_tensor = eq_ind_partial_eval::<F>(suffix);
197
198		let partially_folded_witness = crate::fold_word::BitAxisFolder::new(prefix_tensor.as_ref())
199			.fold::<F, _>(&GlobalAllocator, words);
200
201		inner_product_buffers(&partially_folded_witness, &suffix_tensor)
202	}
203
204	/// Split a GHASH-field element into its `(lo, hi)` 64-bit word pair.
205	fn to_word_pair(elem: F) -> (Word, Word) {
206		let value = u128::from(elem);
207		(Word::from_u64(value as u64), Word::from_u64((value >> 64) as u64))
208	}
209
210	/// Build a valid BinMul witness with `n` random constraints `a * b = c`, where `c` is computed
211	/// with GHASH-field multiplication directly (an independent oracle).
212	fn random_witness(rng: &mut impl Rng, n: usize) -> WordColumns {
213		let mut a_lo = Vec::with_capacity(n);
214		let mut a_hi = Vec::with_capacity(n);
215		let mut b_lo = Vec::with_capacity(n);
216		let mut b_hi = Vec::with_capacity(n);
217		let mut c_lo = Vec::with_capacity(n);
218		let mut c_hi = Vec::with_capacity(n);
219
220		for _ in 0..n {
221			let a = F::random(&mut *rng);
222			let b = F::random(&mut *rng);
223			let c = a * b;
224
225			let (a_lo_i, a_hi_i) = to_word_pair(a);
226			let (b_lo_i, b_hi_i) = to_word_pair(b);
227			let (c_lo_i, c_hi_i) = to_word_pair(c);
228
229			a_lo.push(a_lo_i);
230			a_hi.push(a_hi_i);
231			b_lo.push(b_lo_i);
232			b_hi.push(b_hi_i);
233			c_lo.push(c_lo_i);
234			c_hi.push(c_hi_i);
235		}
236
237		(a_lo, a_hi, b_lo, b_hi, c_lo, c_hi)
238	}
239
240	#[test]
241	fn prove_and_verify() {
242		let mut rng = StdRng::seed_from_u64(0);
243
244		const LOG_N: usize = 5;
245		let (a_lo, a_hi, b_lo, b_hi, c_lo, c_hi) = random_witness(&mut rng, 1 << LOG_N);
246
247		// BinMul commits no oracles.
248		let oracle_specs: [OracleSpec; 0] = [];
249
250		// Run prover.
251		let mut prover_transcript = ProverTranscript::<StdChallenger>::default();
252		let mut prover_channel =
253			NaiveProverChannel::<F, _>::new(&mut prover_transcript, oracle_specs.to_vec());
254		let prove_output = prove::<_, F, P, _>(
255			[&a_lo, &a_hi, &b_lo, &b_hi, &c_lo, &c_hi],
256			&mut prover_channel,
257			&GlobalAllocator,
258		);
259		prover_channel.finish();
260
261		let BinMulOutput {
262			eval_point,
263			a_lo_evals,
264			a_hi_evals,
265			b_lo_evals,
266			b_hi_evals,
267			c_lo_evals,
268			c_hi_evals,
269		} = prove_output.clone();
270
271		// Independently check each column's per-bit evals against the witness MLE. We batch the bit
272		// columns with a `z_challenge` and compare at a single point
273		// `consistency_check_eval_point`.
274		let z_challenge: Vec<F> = (0..Word::LOG_BITS).map(|_| F::random(&mut rng)).collect();
275		let z_tensor = eq_ind_partial_eval::<F>(&z_challenge);
276		let consistency_check_eval_point = [z_challenge, eval_point].concat();
277		let get_consistency_check_eval =
278			|evals: [F; Word::BITS]| izip!(evals, z_tensor.as_ref()).map(|(x, y)| x * y).sum();
279
280		let test_cases = [
281			(&a_lo, a_lo_evals),
282			(&a_hi, a_hi_evals),
283			(&b_lo, b_lo_evals),
284			(&b_hi, b_hi_evals),
285			(&c_lo, c_lo_evals),
286			(&c_hi, c_hi_evals),
287		];
288		for (words, evals) in test_cases {
289			let expected_eval = evaluate_witness(words, &consistency_check_eval_point);
290			let given_eval = get_consistency_check_eval(evals);
291			assert_eq!(expected_eval, given_eval);
292		}
293
294		// Run verifier.
295		let mut verifier_transcript = prover_transcript.into_verifier();
296		let mut verifier_channel =
297			NaiveVerifierChannel::<F, _>::new(&mut verifier_transcript, &oracle_specs);
298		let verify_output = verify(LOG_N, &mut verifier_channel).unwrap();
299		verifier_channel.finish();
300
301		assert_eq!(prove_output, verify_output);
302	}
303
304	// A witness whose constraint count is not a power of two proves exactly as the same witness
305	// zero-padded up to one: byte-identical transcript, identical output claims, and the proof
306	// still verifies against the padded variable count the verifier derives from the constraint
307	// system.
308	//
309	// This is the guarantee that lets the M4 prover stop materializing the padding rows
310	// (BINIUS-388) without moving anything on the wire.
311	#[test]
312	fn unpadded_columns_prove_identically_to_padded_ones() {
313		let mut rng = StdRng::seed_from_u64(2);
314
315		// Constraint counts crossing the fold's 64-word chunk and the packed width: a single row, a
316		// sub-chunk count, a count straddling the chunk boundary, and a multi-chunk count.
317		for n in [1, 3, 65, 100] {
318			let columns = random_witness(&mut rng, n);
319			let n_vars = log2_ceil_usize(n);
320
321			// The same witness with explicit zero rows up to the hypercube's size. A zero row
322			// satisfies `0 * 0 = 0`, so the padded witness is equally valid.
323			let padded = {
324				let pad = |words: &Vec<Word>| {
325					let mut padded = words.clone();
326					padded.resize(1 << n_vars, Word::ZERO);
327					padded
328				};
329				let (a_lo, a_hi, b_lo, b_hi, c_lo, c_hi) = &columns;
330				(pad(a_lo), pad(a_hi), pad(b_lo), pad(b_hi), pad(c_lo), pad(c_hi))
331			};
332
333			// One proof per shape, each on a fresh transcript over the same challenger seed.
334			let run = |columns: &WordColumns| {
335				let (a_lo, a_hi, b_lo, b_hi, c_lo, c_hi) = columns;
336				let oracle_specs: [OracleSpec; 0] = [];
337				let mut transcript = ProverTranscript::<StdChallenger>::default();
338				let mut channel =
339					NaiveProverChannel::<F, _>::new(&mut transcript, oracle_specs.to_vec());
340				let output = prove::<_, F, P, _>(
341					[a_lo, a_hi, b_lo, b_hi, c_lo, c_hi],
342					&mut channel,
343					&GlobalAllocator,
344				);
345				channel.finish();
346				(output, transcript.finalize())
347			};
348			let (unpadded_output, unpadded_proof) = run(&columns);
349			let (padded_output, padded_proof) = run(&padded);
350
351			assert_eq!(unpadded_output, padded_output, "claims differ at n = {n}");
352			assert_eq!(unpadded_proof, padded_proof, "transcript differs at n = {n}");
353
354			let oracle_specs: [OracleSpec; 0] = [];
355			let mut verifier_transcript =
356				VerifierTranscript::new(StdChallenger::default(), unpadded_proof);
357			let mut verifier_channel =
358				NaiveVerifierChannel::<F, _>::new(&mut verifier_transcript, &oracle_specs);
359			let verify_output = verify(n_vars, &mut verifier_channel).unwrap();
360			verifier_channel.finish();
361
362			assert_eq!(unpadded_output, verify_output, "verifier disagrees at n = {n}");
363		}
364	}
365
366	#[test]
367	fn verify_rejects_tampered_c() {
368		let mut rng = StdRng::seed_from_u64(1);
369
370		const LOG_N: usize = 5;
371		let (a_lo, a_hi, b_lo, b_hi, mut c_lo, c_hi) = random_witness(&mut rng, 1 << LOG_N);
372
373		// Corrupt one c_lo word so the constraint no longer holds.
374		c_lo[3] = Word::from_u64(c_lo[3].as_u64() ^ 1);
375
376		let oracle_specs: [OracleSpec; 0] = [];
377
378		let mut prover_transcript = ProverTranscript::<StdChallenger>::default();
379		let mut prover_channel =
380			NaiveProverChannel::<F, _>::new(&mut prover_transcript, oracle_specs.to_vec());
381		let _ = prove::<_, F, P, _>(
382			[&a_lo, &a_hi, &b_lo, &b_hi, &c_lo, &c_hi],
383			&mut prover_channel,
384			&GlobalAllocator,
385		);
386		prover_channel.finish();
387
388		let mut verifier_transcript = prover_transcript.into_verifier();
389		let mut verifier_channel =
390			NaiveVerifierChannel::<F, _>::new(&mut verifier_transcript, &oracle_specs);
391		let result = verify(LOG_N, &mut verifier_channel);
392		assert!(result.is_err(), "verifier must reject a tampered witness");
393	}
394}