1use std::{iter, mem, ops::Deref};
5
6use binius_field::{BinaryField, Field, PackedField};
7use binius_iop::fri::{FRIParams, fold::fold_chunk};
8use binius_math::{
9 FieldBuffer, FieldSlice, inner_product::inner_product_buffers,
10 multilinear::eq::eq_ind_partial_eval, ntt::AdditiveNTT,
11};
12use binius_utils::{checked_arithmetics::log2_ceil_usize, rayon::prelude::*};
13use tracing::instrument;
14
15use super::query::FRIQueryProver;
16use crate::{
17 fri::{BatchBrakedownOracleProver, BrakedownOracleProver, FRIOracleProver},
18 merkle_channel::MerkleIPProverChannel,
19};
20
21enum FRIFolderState<P, C, Data = Vec<P>>
22where
23 P: PackedField,
24 Data: Deref<Target = [P]>,
25{
26 FirstFold(BatchBrakedownFolder<P, C, Data>),
27 LaterFolds {
28 first_oracle: BatchBrakedownOracleProver<P, C, Data>,
29 last_codeword: FieldBuffer<P::Scalar>,
30 last_commitment: C,
31 round_oracles: Vec<FRIOracleProver<P::Scalar, C>>,
32 },
33}
34pub struct FRIFoldProver<'a, F, P, NTT, C, Data = Vec<P>>
39where
40 F: BinaryField,
41 P: PackedField<Scalar = F>,
42 Data: Deref<Target = [P]>,
43{
44 params: &'a FRIParams<F>,
45 ntt: &'a NTT,
46 state: Option<FRIFolderState<P, C, Data>>,
47 curr_round: usize,
48 next_commit_round: Option<usize>,
49 unprocessed_challenges: Vec<F>,
50}
51
52impl<'a, F, P, NTT, C, Data> FRIFoldProver<'a, F, P, NTT, C, Data>
53where
54 F: BinaryField,
55 P: PackedField<Scalar = F>,
56 NTT: AdditiveNTT<Field = F> + Sync,
57 Data: Deref<Target = [P]>,
58{
59 pub fn new(
61 params: &'a FRIParams<F>,
62 ntt: &'a NTT,
63 committed_codeword: FieldBuffer<P, Data>,
64 commitment: C,
65 ) -> Self {
66 Self::new_batch(params, ntt, vec![(committed_codeword, commitment)])
67 }
68
69 pub fn new_batch(
85 params: &'a FRIParams<F>,
86 ntt: &'a NTT,
87 committed_codewords: Vec<(FieldBuffer<P, Data>, C)>,
88 ) -> Self {
89 let input_oracles = params.input_oracles();
90 assert_eq!(
91 committed_codewords.len(),
92 input_oracles.len(),
93 "precondition: committed_codewords.len() must equal params.input_oracles().len()"
94 );
95
96 let log_dim = params.rs_code().log_dim();
101 let log_inv_rate = params.rs_code().log_inv_rate();
102 for spec in input_oracles {
103 assert!(
104 spec.log_lift <= log_dim,
105 "precondition: input oracle dimension must not exceed the reduced code dimension"
106 );
107 }
108
109 let folders = iter::zip(committed_codewords, input_oracles)
110 .map(|((codeword, commitment), spec)| {
111 let oracle_log_dim = log_dim - spec.log_lift;
115 let expected_log_len = oracle_log_dim + spec.log_batch_size() + log_inv_rate;
116 assert_eq!(
117 codeword.log_len(),
118 expected_log_len,
119 "precondition: interleaved codeword length must match the oracle's \
120 Reed-Solomon code length plus its batch size"
121 );
122 ProxTestFolder {
123 log_early_batch_size: spec.log_early_batch_size,
124 log_later_batch_size: spec.log_later_batch_size,
125 log_lift: spec.log_lift,
126 codeword,
127 commitment,
128 }
129 })
130 .collect::<Vec<_>>();
131 let batch_folder = BatchBrakedownFolder::new(folders, params.rs_code().log_len());
132
133 let next_commit_round = Some(params.log_batch_size());
134 Self {
135 params,
136 ntt,
137 state: Some(FRIFolderState::FirstFold(batch_folder)),
138 curr_round: 0,
139 next_commit_round,
140 unprocessed_challenges: Vec::with_capacity(params.rs_code().log_dim()),
141 }
142 }
143
144 pub const fn n_rounds(&self) -> usize {
146 self.params.n_fold_rounds()
147 }
148
149 fn is_commitment_round(&self) -> bool {
150 self.next_commit_round
151 .is_some_and(|round| round == self.curr_round)
152 }
153
154 pub fn receive_challenge(&mut self, challenge: F) {
171 self.unprocessed_challenges.push(challenge);
172 self.curr_round += 1;
173 }
174
175 pub fn execute_fold_round<Channel>(&mut self, channel: &mut Channel)
187 where
188 Channel: MerkleIPProverChannel<F, Commitment = C>,
189 {
190 if !self.is_commitment_round() {
191 return;
192 }
193
194 let state = self
195 .state
196 .take()
197 .expect("state is always Some by struct invariant");
198
199 let new_state = match state {
200 FRIFolderState::FirstFold(folder) => {
201 let _scope = tracing::debug_span!(
202 "FRI Initial Fold",
203 log_len = folder.log_len(),
204 arity = self.unprocessed_challenges.len()
205 )
206 .entered();
207
208 let challenges = mem::take(&mut self.unprocessed_challenges);
212 let (folded_codeword, first_oracle) = folder.fold(&challenges);
213
214 let next_arity = self.params.fold_arities().first().copied();
215 let last_commitment = self.commit_round(channel, &folded_codeword, next_arity);
216
217 FRIFolderState::LaterFolds {
218 first_oracle,
219 last_codeword: folded_codeword,
220 last_commitment,
221 round_oracles: Vec::with_capacity(self.params.fold_arities().len()),
222 }
223 }
224 FRIFolderState::LaterFolds {
225 first_oracle,
226 last_codeword,
227 last_commitment,
228 mut round_oracles,
229 } => {
230 let _fri_round_scope = tracing::debug_span!(
231 "FRI Round Fold",
232 log_len = last_codeword.log_len(),
233 arity = self.unprocessed_challenges.len()
234 )
235 .entered();
236
237 let challenges = mem::take(&mut self.unprocessed_challenges);
240 let fri_fold_span = tracing::debug_span!("FRI Fold").entered();
241 let folded_codeword = fold_codeword(self.ntt, last_codeword.as_view(), &challenges);
242 drop(fri_fold_span);
243 let oracle = FRIOracleProver::new(last_codeword, last_commitment, challenges.len());
246
247 let next_arity = self
248 .params
249 .fold_arities()
250 .get(round_oracles.len() + 1)
251 .copied();
252 let last_commitment = self.commit_round(channel, &folded_codeword, next_arity);
253
254 round_oracles.push(oracle);
255 FRIFolderState::LaterFolds {
256 first_oracle,
257 last_codeword: folded_codeword,
258 last_commitment,
259 round_oracles,
260 }
261 }
262 };
263
264 self.state = Some(new_state);
265 }
266
267 fn commit_round<Channel>(
275 &mut self,
276 channel: &mut Channel,
277 folded_codeword: &FieldBuffer<F>,
278 next_arity: Option<usize>,
279 ) -> C
280 where
281 Channel: MerkleIPProverChannel<F, Commitment = C>,
282 {
283 let log_coset_size = next_arity.unwrap_or_else(|| self.params.n_final_challenges());
284
285 let _merkle_tree_span = tracing::debug_span!("Merkle Tree").entered();
286 let commitment =
287 channel.send_merkle_commitment(folded_codeword.as_view(), 1 << log_coset_size);
288
289 self.next_commit_round = next_arity.map(|arity| self.curr_round + arity);
292
293 commitment
294 }
295
296 #[instrument(skip_all, name = "fri::FRIFolder::finalize", level = "debug")]
309 #[allow(clippy::type_complexity)]
310 pub fn finalize(mut self) -> (FieldBuffer<F>, C, FRIQueryProver<F, P, C, Data>) {
311 assert_eq!(
312 self.curr_round,
313 self.n_rounds(),
314 "precondition: all fold rounds must be executed before finalize"
315 );
316
317 self.unprocessed_challenges.clear();
318
319 match self
320 .state
321 .take()
322 .expect("state is always Some by struct invariant")
323 {
324 FRIFolderState::LaterFolds {
328 first_oracle,
329 last_codeword,
330 last_commitment,
331 round_oracles,
332 } => {
333 let query_prover = FRIQueryProver::new(first_oracle, round_oracles);
334 (last_codeword, last_commitment, query_prover)
335 }
336 FRIFolderState::FirstFold(_) => {
339 unreachable!("the first fold always runs before curr_round reaches n_rounds")
340 }
341 }
342 }
343
344 pub fn finish_proof<Channel>(self, channel: &mut Channel)
353 where
354 Channel: MerkleIPProverChannel<F, Commitment = C>,
355 {
356 let n_test_queries = self.params.n_test_queries();
357 let index_bits = self.params.index_bits();
358 let (terminate_codeword, terminal_commitment, query_prover) = self.finalize();
359
360 let indices = (0..n_test_queries)
364 .map(|_| channel.sample_bits(index_bits))
365 .collect::<Vec<_>>();
366
367 query_prover.prove_queries(&indices, channel);
369 channel.send_committed_vector(&terminal_commitment, terminate_codeword.as_view());
370 }
371}
372
373#[instrument(skip_all, level = "debug")]
386fn fold_codeword<F, NTT>(ntt: &NTT, codeword: FieldSlice<'_, F>, challenges: &[F]) -> FieldBuffer<F>
387where
388 F: BinaryField,
389 NTT: AdditiveNTT<Field = F> + Sync,
390{
391 let log_len = codeword.log_len();
392 assert!(challenges.len() <= log_len);
393
394 let folded_log_len = log_len - challenges.len();
395
396 let chunk_size = 1 << challenges.len();
398 let values: Vec<F> = codeword
399 .par_chunks(challenges.len())
400 .enumerate()
401 .map_init(
402 || vec![F::default(); chunk_size],
403 |scratch_buffer, (i, chunk)| {
404 scratch_buffer.copy_from_slice(chunk.as_ref());
405 fold_chunk(ntt, log_len, i, scratch_buffer, challenges)
406 },
407 )
408 .collect();
409 FieldBuffer::new(folded_log_len, values)
410}
411
412pub struct ProxTestFolder<P: PackedField, C, Data: Deref<Target = [P]> = Vec<P>> {
413 log_early_batch_size: usize,
417 log_later_batch_size: usize,
421 log_lift: usize,
424 codeword: FieldBuffer<P, Data>,
425 commitment: C,
426}
427
428impl<P: PackedField, C, Data: Deref<Target = [P]>> ProxTestFolder<P, C, Data> {
429 const fn log_batch_size(&self) -> usize {
431 self.log_early_batch_size + self.log_later_batch_size
432 }
433
434 pub const fn log_folded_len(&self) -> usize {
435 self.codeword.log_len() - self.log_batch_size()
436 }
437}
438
439pub struct BatchBrakedownFolder<P: PackedField, C, Data: Deref<Target = [P]> = Vec<P>> {
446 log_code_len: usize,
447 folders: Vec<ProxTestFolder<P, C, Data>>,
448}
449
450impl<F: Field, P: PackedField<Scalar = F>, C, Data: Deref<Target = [P]>>
451 BatchBrakedownFolder<P, C, Data>
452{
453 pub fn new(folders: Vec<ProxTestFolder<P, C, Data>>, log_code_len: usize) -> Self {
459 assert!(!folders.is_empty()); for folder in &folders {
461 assert!(folder.log_folded_len() <= log_code_len);
462 }
463 Self {
464 log_code_len,
465 folders,
466 }
467 }
468
469 fn log_len(&self) -> usize {
471 let max_folder_log_len = self
472 .folders
473 .iter()
474 .map(|folder| folder.codeword.log_len())
475 .max()
476 .expect("folders is not empty by struct invariant");
477 let log_folders = log2_ceil_usize(self.folders.len());
478 max_folder_log_len + log_folders
479 }
480
481 pub fn fold(
482 self,
483 challenges: &[F],
484 ) -> (FieldBuffer<F>, BatchBrakedownOracleProver<P, C, Data>) {
485 let max_early = self
489 .folders
490 .iter()
491 .map(|folder| folder.log_early_batch_size)
492 .max()
493 .expect("folders is not empty by struct invariant");
494 let max_later = self
495 .folders
496 .iter()
497 .map(|folder| folder.log_later_batch_size)
498 .max()
499 .expect("folders is not empty by struct invariant");
500 let log_n_oracles = log2_ceil_usize(self.folders.len());
501
502 let early_challenges = &challenges[..max_early];
503 let outer_challenges = &challenges[max_early..max_early + log_n_oracles];
504 let later_challenges = &challenges[max_early + log_n_oracles..];
505 let outer_tensor = eq_ind_partial_eval::<F>(outer_challenges);
506
507 let code_len = 1 << self.log_code_len;
512 let mut values = Vec::<F>::with_capacity(code_len);
513 let mut oracles = Vec::with_capacity(self.folders.len());
514
515 for (index, (folder, &scalar)) in iter::zip(self.folders, outer_tensor.as_ref()).enumerate()
516 {
517 let ProxTestFolder {
518 log_early_batch_size,
519 log_later_batch_size,
520 log_lift,
521 codeword,
522 commitment,
523 } = folder;
524 let log_batch_size = log_early_batch_size + log_later_batch_size;
525
526 let early_window = &early_challenges[max_early - log_early_batch_size..];
530 let later_window = &later_challenges[max_later - log_later_batch_size..];
531 let fold_challenges: Vec<F> =
532 early_window.iter().chain(later_window).copied().collect();
533
534 let mut tensor = eq_ind_partial_eval::<P>(&fold_challenges);
540 if scalar != F::ONE {
541 let scalar_broadcast = P::broadcast(scalar);
542 for packed in tensor.as_mut() {
543 *packed *= scalar_broadcast;
544 }
545 }
546
547 if index == 0 {
553 values.spare_capacity_mut()[..code_len]
554 .par_chunks_mut(1 << log_lift)
555 .zip(codeword.par_chunks(log_batch_size))
556 .for_each(|(copies, chunk)| {
557 let value = inner_product_buffers(&chunk, &tensor);
558 for acc in copies {
559 acc.write(value);
560 }
561 });
562 debug_assert_eq!(
565 code_len >> log_lift,
566 1 << (codeword.log_len() - log_batch_size),
567 "the folded codeword must hold one entry per lifted output chunk"
568 );
569 unsafe { values.set_len(code_len) };
572 } else {
573 values
574 .par_chunks_mut(1 << log_lift)
575 .zip(codeword.par_chunks(log_batch_size))
576 .for_each(|(copies, chunk)| {
577 let value = inner_product_buffers(&chunk, &tensor);
578 for acc in copies {
579 *acc += value;
580 }
581 });
582 }
583
584 oracles.push(BrakedownOracleProver::new(codeword, commitment, log_lift));
585 }
586
587 let combined_codeword = FieldBuffer::new(self.log_code_len, values);
588 (combined_codeword, BatchBrakedownOracleProver::new(oracles))
589 }
590}
591
592#[cfg(test)]
593mod tests {
594 use binius_field::Ghash128b as B128;
595 use binius_math::{
596 BinarySubspace,
597 ntt::{NeighborsLastReference, domain_context::GenericOnTheFly},
598 test_utils::{random_field_buffer, random_scalars},
599 };
600 use proptest::prelude::*;
601 use rand::prelude::*;
602
603 use super::*;
604
605 proptest! {
606 #[test]
607 fn test_fri_compatible_ntt_domains(log_dim in 0..8usize, arity in 0..4usize) {
608 test_help_fri_compatible_ntt_domains(log_dim, arity);
609 }
610 }
611
612 fn test_help_fri_compatible_ntt_domains(log_dim: usize, arity: usize) {
613 let subspace = BinarySubspace::with_dim(32);
614 let domain_context = GenericOnTheFly::generate_from_subspace(&subspace);
615 let ntt = NeighborsLastReference { domain_context };
616
617 let mut rng = StdRng::seed_from_u64(0);
618 let msg = random_field_buffer(&mut rng, log_dim + arity);
619 let challenges = random_scalars(&mut rng, arity);
620
621 let query = eq_ind_partial_eval::<B128>(&challenges);
622
623 let folded_vals: Vec<B128> = msg
626 .chunks(arity)
627 .map(|row| inner_product_buffers(&row, &query))
628 .collect();
629 let mut folded_msg = FieldBuffer::new(log_dim, folded_vals);
630 assert_eq!(folded_msg.log_len(), log_dim);
631
632 let mut codeword = msg;
634 ntt.forward_transform(codeword.as_mut_view(), 0, 0);
635
636 let folded_codeword = fold_codeword(&ntt, codeword.as_view(), &challenges);
638
639 ntt.forward_transform(folded_msg.as_mut_view(), 0, 0);
641
642 assert_eq!(folded_codeword, folded_msg);
644 }
645}