binius_prover/protocols/bitand/
prover.rs1use std::ops::Deref;
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, PackedField, Rijndael8b as B8};
9use binius_ip_prover::sumcheck::{common::MleCheckProver, quadratic_mlecheck_prover};
10use binius_math::{BinarySubspace, univariate::EvaluationDomain};
11use binius_verifier::{
12 config::PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES, protocols::bitand::ROWS_PER_HYPERCUBE_VERTEX,
13};
14
15use super::sumcheck_round_messages;
16use crate::fold_word::BitAxisFolder;
17
18pub struct UnivariateRoundProver<F, Data>
29where
30 F: BinaryField,
31{
32 log_words: usize,
33 first_col: Data,
34 second_col: Data,
35 big_field_zerocheck_challenges: Vec<F>,
36 univariate_round_message: [F; ROWS_PER_HYPERCUBE_VERTEX],
37 univariate_round_message_domain: BinarySubspace<F>,
38}
39
40impl<F, Data> UnivariateRoundProver<F, Data>
41where
42 F: BinaryField + From<B8>,
43 Data: Deref<Target = [Word]>,
44{
45 pub fn compute_message(
82 log_words: usize,
83 first_col: Data,
84 second_col: Data,
85 big_field_zerocheck_challenges: Vec<F>,
86 prover_message_domain: &BinarySubspace<B8>,
87 ) -> Self {
88 let univariate_round_message = tracing::debug_span!("Compute univariate round message")
89 .in_scope(|| {
90 sumcheck_round_messages::univariate_round_message_extension_domain::<F>(
91 log_words,
92 &first_col,
93 &second_col,
94 &big_field_zerocheck_challenges,
95 prover_message_domain,
96 )
97 });
98
99 Self {
100 log_words,
101 first_col,
102 second_col,
103 univariate_round_message,
104 big_field_zerocheck_challenges,
105 univariate_round_message_domain: prover_message_domain.isomorphic(),
106 }
107 }
108
109 pub const fn round_message(&self) -> &[F; ROWS_PER_HYPERCUBE_VERTEX] {
115 &self.univariate_round_message
116 }
117
118 pub fn univariate_claim(&self, challenge: F) -> F {
123 let mut coeffs = vec![F::ZERO; 2 * ROWS_PER_HYPERCUBE_VERTEX];
124 coeffs[ROWS_PER_HYPERCUBE_VERTEX..].copy_from_slice(&self.univariate_round_message);
125 self.univariate_round_message_domain
126 .extrapolate(&coeffs, &challenge)
127 }
128
129 pub fn fold<'alloc, PChallenge: PackedField<Scalar = F>, A: Allocator>(
147 self,
148 alloc: &'alloc A,
149 challenge: F,
150 ) -> impl MleCheckProver<F> + 'alloc {
151 let claim = self.univariate_claim(challenge);
152 let round_message_domain = &self.univariate_round_message_domain;
153 let univariate_domain = round_message_domain.reduce_dim(round_message_domain.dim() - 1);
154 let lagrange_evals = univariate_domain.lagrange_evals(&challenge);
155 let folder = BitAxisFolder::new(&lagrange_evals);
156
157 let proving_polys =
158 folder.fold_bitand_operands::<PChallenge, _>(alloc, &self.first_col, &self.second_col);
159
160 let upcasted_small_field_challenges = PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES
161 .iter()
162 .copied()
163 .take(self.log_words)
164 .map(F::from);
165
166 let verifier_field_zerocheck_challenges = upcasted_small_field_challenges
167 .chain(self.big_field_zerocheck_challenges)
168 .collect::<Vec<_>>();
169
170 quadratic_mlecheck_prover(
171 alloc,
172 proving_polys,
173 |[a, b, c]| a * b - c,
174 |[a, b, _]| a * b,
175 verifier_field_zerocheck_challenges,
176 claim,
177 )
178 }
179}
180
181#[cfg(test)]
182mod test {
183 use std::{iter, iter::repeat_with};
184
185 use binius_compute::GlobalAllocator;
186 use binius_core::word::Word;
187 use binius_field::{Rijndael8b, arch::OptimalPackedB128};
188 use binius_math::{
189 BinarySubspace, FieldBuffer, multilinear::evaluate::evaluate, univariate::EvaluationDomain,
190 };
191 use binius_transcript::ProverTranscript;
192 use binius_verifier::{
193 config::{B128, StdChallenger},
194 protocols::bitand::{AndCheckOutput, SKIPPED_VARS},
195 verify_bitand_reduction,
196 };
197 use rand::prelude::*;
198
199 use crate::{fold_word::BitAxisFolder, protocols::bitand::prove};
200
201 fn random_words(log_num_words: usize, mut rng: impl Rng) -> Vec<Word> {
202 repeat_with(|| Word(rng.random()))
203 .take(1 << log_num_words)
204 .collect()
205 }
206
207 #[test]
208 fn test_transcript_prover_verifies() {
209 let mut prover_challenger = ProverTranscript::new(StdChallenger::default());
210 let log_num_rows = 6;
211 let mut rng = StdRng::seed_from_u64(0);
212
213 let first_mlv = random_words(log_num_rows, &mut rng);
214 let second_mlv = random_words(log_num_rows, &mut rng);
215 let third_mlv: Vec<Word> = iter::zip(&first_mlv, &second_mlv)
218 .map(|(&a, &b)| a & b)
219 .collect();
220
221 let prover_message_domain = BinarySubspace::<Rijndael8b>::with_dim(SKIPPED_VARS + 1);
223 let verifier_message_domain = prover_message_domain.isomorphic();
224
225 let prove_output = prove::<_, B128, OptimalPackedB128, _, _>(
226 [first_mlv.clone(), second_mlv.clone()],
227 &[],
228 &mut prover_challenger,
229 &GlobalAllocator,
230 );
231
232 let mut verifier_challenger = prover_challenger.into_verifier();
233 let verify_output = verify_bitand_reduction(
234 log_num_rows,
235 &verifier_message_domain,
236 &[],
237 &mut verifier_challenger,
238 )
239 .unwrap();
240
241 assert_eq!(prove_output, verify_output);
242
243 let AndCheckOutput {
244 z_challenge,
245 rerand,
246 } = verify_output;
247 let [a_eval, b_eval, c_eval] = rerand.bitand_evals;
248 let eval_point = rerand.eval_point;
249
250 let verifier_univariate_domain = verifier_message_domain.reduce_dim(SKIPPED_VARS);
251
252 let one_bit_mlvs = [first_mlv, second_mlv, third_mlv];
253
254 let verifier_lagrange_evals = verifier_univariate_domain.lagrange_evals(&z_challenge);
255 let folder = BitAxisFolder::new(&verifier_lagrange_evals);
256 for (i, eval) in [a_eval, b_eval, c_eval].iter().enumerate() {
257 let folded: FieldBuffer<B128> = folder.fold(&GlobalAllocator, &one_bit_mlvs[i]);
258 assert_eq!(evaluate(&folded, &eval_point), *eval);
259 }
260 }
261}