Skip to main content

binius_iop_prover/fri/
fold.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use 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}
34/// A stateful prover for the FRI fold phase.
35///
36/// Fold-round codewords are committed by sending them over a Merkle channel with commitment
37/// handle type `C`, matching the channel's `Commitment` associated type.
38pub 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	/// Constructs a new folder for a single committed input oracle.
60	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	/// Constructs a new folder for a batch of committed input oracles.
70	///
71	/// The input oracles share the Reed-Solomon code but may have differing batch sizes; they are
72	/// folded and combined into a single first-round codeword. The codewords must be supplied in
73	/// the same order as [`FRIParams::input_oracles`], each with the commitment handle produced
74	/// when it was sent over the Merkle channel.
75	///
76	/// ## Preconditions
77	///
78	/// * `committed_codewords.len()` must equal `params.input_oracles().len()`.
79	/// * Each input oracle's dimension (`rs_code().log_dim() - log_lift`) must be at most
80	///   `params.rs_code().log_dim()`.
81	/// * Each codeword's length must equal its oracle's Reed-Solomon code length plus its batch
82	///   size (`rs_code().log_dim() - log_lift + log_batch_size + log_inv_rate`), and its
83	///   commitment's leaf size must be one interleaved coset (`2^log_batch_size` scalars).
84	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		// Each input oracle's Reed-Solomon dimension (`log_dim - log_lift`) must not exceed the
97		// first-round (reduced) code dimension; smaller oracles are lifted (padded) to it. This
98		// holds whenever `log_lift <= log_dim`, so assert it here rather than trusting the
99		// caller.
100		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				// The oracle's own codeword has dimension `log_dim - log_lift`, so its interleaved
112				// length is that plus the batch size plus the inverse rate. It is lifted to the
113				// common first-round length by duplicating each entry `2^log_lift` times.
114				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	/// Number of fold rounds, including the final fold.
145	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	/// Records the folding challenge for the current round and advances the round counter.
155	///
156	/// The challenge is buffered, not applied immediately.
157	/// Buffered challenges are consumed lazily at the next commit round.
158	/// This avoids materializing intermediate folded codewords.
159	///
160	/// The challenge order is a hard contract, shared with the verifier's fold order.
161	/// Feed challenges in the protocol's round order:
162	///
163	/// 1. the shared mask challenge `gamma`, once, if any oracle is ZK (the inner unbatch round);
164	/// 2. the `log_n_oracles` outer batching challenges (the oracle-combine rounds);
165	/// 3. one challenge per MLE-check round over the combined oracle's variables.
166	///
167	/// Steps 1 and 2 are fed up front.
168	/// Step 3 is interleaved with the fold rounds: one challenge after each fold.
169	/// The total number of challenges must equal the number of fold rounds.
170	pub fn receive_challenge(&mut self, challenge: F) {
171		self.unprocessed_challenges.push(challenge);
172		self.curr_round += 1;
173	}
174
175	/// Executes the next fold round, committing the folded codeword over the channel if this is a
176	/// commitment round.
177	///
178	/// On a commitment round, the folded codeword's Merkle commitment is computed and its root is
179	/// written to the channel as an observed message. Call this *after* writing any other messages
180	/// belonging to the same round (e.g. sumcheck round coefficients), so the root lands after them
181	/// in the transcript.
182	///
183	/// As a memory efficient optimization, this method may not actually do the folding, but instead
184	/// accumulate the folding challenge for processing at a later time. This saves us from storing
185	/// intermediate folded codewords.
186	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				// Fold the batch of interleaved codewords that were originally committed into a
209				// single codeword with the same block length, and turn them into a batched
210				// Brakedown query oracle.
211				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				// Fold a full codeword committed in the previous FRI round into a codeword with
238				// reduced dimension and rate.
239				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				// The fold consuming `last_codeword` has arity `challenges.len()`, which is the
244				// coset size its commitment was built with.
245				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	/// Commits a folded codeword over the channel, advancing `next_commit_round` for the next
268	/// fold.
269	///
270	/// The coset (leaf) size is determined by the arity of the *next* fold round, or by the number
271	/// of final challenges once there are no more committed rounds (the terminal codeword). The
272	/// returned commitment handle owns the committed tree, so the query phase can open it over the
273	/// channel later.
274	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		// The next commitment lands `next_arity` rounds after the current one. Once there is no
290		// next arity, this is the terminal codeword and no further commitments are made.
291		self.next_commit_round = next_arity.map(|arity| self.curr_round + arity);
292
293		commitment
294	}
295
296	/// Finalizes the FRI folding process.
297	///
298	/// This step will process any unprocessed folding challenges to produce the
299	/// final folded codeword. Then it will decode this final folded codeword
300	/// to get the final message.
301	///
302	/// This returns the terminal codeword, its commitment handle (for sending it in full over a
303	/// Merkle channel), and a query prover instance.
304	///
305	/// ## Preconditions
306	///
307	/// * All fold rounds must have been executed (`curr_round == n_rounds()`).
308	#[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			// The final fold round produced the terminal codeword and committed the prior one, so
325			// the state is always `LaterFolds` once `curr_round` reaches `n_rounds`. The
326			// terminal codeword is sent in full and therefore is not wrapped in a query oracle.
327			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			// The first fold fires at `curr_round == log_batch_size <= n_rounds` and
337			// `execute_fold_round` runs every round, so the first fold always precedes `finalize`.
338			FRIFolderState::FirstFold(_) => {
339				unreachable!("the first fold always runs before curr_round reaches n_rounds")
340			}
341		}
342	}
343
344	/// Runs the FRI query phase over the channel.
345	///
346	/// Samples the query indices, sends the per-oracle batched query openings, and sends the
347	/// terminal codeword in full.
348	///
349	/// ## Preconditions
350	///
351	/// * All fold rounds must have been executed (`curr_round == n_rounds()`).
352	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		// Sample all query indices before sending the (per-oracle batched) query openings. The
361		// decommitment advice is not absorbed by the challenger, so this matches the verifier
362		// sampling all indices up front.
363		let indices = (0..n_test_queries)
364			.map(|_| channel.sample_bits(index_bits))
365			.collect::<Vec<_>>();
366
367		// Send the per-oracle batched query openings, then the terminal codeword in full.
368		query_prover.prove_queries(&indices, channel);
369		channel.send_committed_vector(&terminal_commitment, terminate_codeword.as_view());
370	}
371}
372
373/// FRI-fold the codeword using the given challenges.
374///
375/// ## Arguments
376///
377/// * `ntt` - the NTT instance, used to look up the twiddle values.
378/// * `codeword` - an interleaved codeword.
379/// * `challenges` - the folding challenges. The length must be at least `log_batch_size`.
380/// * `log_len` - the binary logarithm of the code length.
381///
382/// See [DP24], Def. 3.6 and Lemma 3.9 for more details.
383///
384/// [DP24]: <https://eprint.iacr.org/2024/504>
385#[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	// For each coset of size `2^chunk_size` in the codeword, fold it with the folding challenges.
397	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	/// log2 the number of *early* batch-fold challenges this oracle's interleaving folds with
414	/// (sampled before the outer oracle-combine challenges). The oracle folds with the
415	/// `log_early_batch_size`-length suffix of the early challenges.
416	log_early_batch_size: usize,
417	/// log2 the number of *later* batch-fold challenges this oracle's interleaving folds with
418	/// (sampled after the outer oracle-combine challenges). The oracle folds with the
419	/// `log_later_batch_size`-length suffix of the later challenges.
420	log_later_batch_size: usize,
421	/// log2 the lift factor (oracle padding): how many times each folded codeword entry is
422	/// duplicated to reach the common first-round length. Zero when no lifting is needed.
423	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	/// The total interleave batch size, `log_early_batch_size + log_later_batch_size`.
430	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
439/// Folds and commits a batch of interleaved codewords that share a folded length.
440///
441/// Each [`ProxTestFolder`] is committed separately and folds (by the same challenges) to a codeword
442/// of the common length `codeword.log_len() - log_batch_size`. The folded codewords are summed into
443/// a single codeword that continues through the FRI rounds, and the per-commitment
444/// [`BrakedownOracleProver`]s are bundled into a [`BatchBrakedownOracleProver`].
445pub 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	/// Constructs a batch folder from one or more interleaved-codeword folders.
454	///
455	/// `log_code_len` is the common (first-round) codeword length the folders combine into. Each
456	/// folder's own folded length must not exceed it; folders that fall short are lifted (their
457	/// folded codewords duplicated) up to `log_code_len` during [`Self::fold`].
458	pub fn new(folders: Vec<ProxTestFolder<P, C, Data>>, log_code_len: usize) -> Self {
459		assert!(!folders.is_empty()); // precondition
460		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	/// Log2 length of the (interleaved) codewords being folded.
470	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		// The first-fold challenge slice is `[early ++ outer ++ later]`: `max_early` early
486		// within-oracle batch challenges, then `log_n_oracles` outer oracle-combine challenges,
487		// then `max_later` later within-oracle batch challenges.
488		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		// The combined codeword is the largest buffer of the fold phase.
508		// It starts uninitialized, and the first oracle writes it rather than adding into it.
509		// Adding to zero is a copy, so a zeroed buffer would cost a fill of the whole buffer.
510		// It would also cost a read of everything that fill had just written.
511		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			// This oracle folds its interleaving with `early_window ++ later_window`, where each
527			// window is the suffix of its group. An oracle is purely early (ZK) or purely later
528			// (non-ZK), so in practice one window is empty, but the concatenation is general.
529			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			// Fold the outer-challenge tensor value into the inner folding tensor so that every
535			// folded entry comes out already scaled by `scalar`. This replaces one scaling mul per
536			// (lifted) output entry with a single pass over the `2^log_batch_size`-element tensor.
537			// A single oracle carries no outer challenges, so its entry is one.
538			// The pass is then a multiply by one at every element, and is skipped.
539			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			// Fold each `2^log_batch_size`-element interleaved chunk into a single scaled value via
548			// an inner product with the (pre-scaled) tensor, landing it directly in the folded
549			// entry's `2^log_lift` contiguous copies in the combined codeword (the Reed-Solomon
550			// codeword duplication identity: `combined[j] += folded[j >> log_lift]`).
551			// No temporary buffer; `1 << log_lift` is `1` when there is no lifting.
552			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				// Pairing two chunk iterators stops at the shorter one.
563				// So every slot is written only when the two sides hold equally many chunks.
564				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				// SAFETY: the chunks partition the slots, and each writes all of its own.
570				// The counts above are equal, so the loop wrote every one of them.
571				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		// Fold the message using regular folding: combine the low `arity` columns of each row
624		// with the eq tensor of the challenges (a partial evaluation of each row at the point).
625		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		// Encode the message over the large domain.
633		let mut codeword = msg;
634		ntt.forward_transform(codeword.as_mut_view(), 0, 0);
635
636		// Fold the encoded message using FRI folding.
637		let folded_codeword = fold_codeword(&ntt, codeword.as_view(), &challenges);
638
639		// Encode the folded message.
640		ntt.forward_transform(folded_msg.as_mut_view(), 0, 0);
641
642		// Check that folding and encoding commute.
643		assert_eq!(folded_codeword, folded_msg);
644	}
645}