1use 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
18fn 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 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 if !lo_tail.is_empty() {
58 values.push(P::from_scalars(iter::zip(lo_tail, hi_tail).map(elem)));
59 }
60
61 values.resize(packed_len, P::zero());
64
65 FieldBuffer::new(n_vars, values)
66}
67
68pub 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 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 let r_z = channel.sample_many(n_vars);
109
110 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 eval_point.reverse();
127
128 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 type WordColumns = (Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>, Vec<Word>);
189
190 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 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 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 let oracle_specs: [OracleSpec; 0] = [];
249
250 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 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 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 #[test]
312 fn unpadded_columns_prove_identically_to_padded_ones() {
313 let mut rng = StdRng::seed_from_u64(2);
314
315 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 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 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 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}