Skip to main content

binius_prover/
ring_switch.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{array, iter, ops::Deref};
5
6use binius_compute::{Allocator, VecLike};
7use binius_core::word::Word;
8use binius_field::{
9	ExtensionField, PackedField,
10	linear_transformation::{
11		BytewiseLookupTransformationFactory, InputWrappingTransformationFactory,
12		LinearTransformationFactory, OutputWrappingTransformationFactory, Transformation,
13	},
14	packed_extension,
15};
16use binius_ip_prover::{
17	channel::IPProverChannel,
18	sumcheck::{bivariate_product_prover, prove_single},
19};
20use binius_math::{
21	FieldBuffer, FieldSlice, FieldVec, inner_product::inner_product_packed,
22	multilinear::eq::eq_ind_partial_eval, tensor_algebra::TensorAlgebra,
23};
24use binius_utils::{checked_arithmetics::log2_ceil_usize, rayon::prelude::*};
25use binius_verifier::{
26	config::{B1, B128, LOG_WORDS_PER_ELEM},
27	protocols::shift::evaluate_words_mle,
28};
29use itertools::izip;
30
31use crate::{
32	bit_matrix::{ColumnSums, LOG_WEIGHTS_PER_TABLE, RowFoldTables},
33	prove::pack_witness,
34};
35
36/// Base-2 log of the low-factor length the tensor split targets.
37///
38/// The row fold consumes a chunk of as many rows as a row has columns.
39/// So the low factor spans exactly that many rows.
40/// For 128 columns that is 16 tables of 256 field elements, or 64 KiB.
41/// That footprint stays resident in a performance core's L1 data cache across every hi-block.
42/// A wider factor needs proportionally more tables and starts missing that cache.
43pub const LOG_SPLIT_BLOCK: usize = <B128 as ExtensionField<B1>>::LOG_DEGREE;
44
45/// Number of subset-sum tables one chunk of the split fold uses.
46const N_ROW_TABLES: usize = 1 << (LOG_SPLIT_BLOCK - LOG_WEIGHTS_PER_TABLE);
47
48/// Folds the 1-bit rows of a matrix against an equality tensor supplied as two factors.
49///
50/// # Overview
51///
52/// The equality tensor over `n_lo + n_hi` coordinates factors along any coordinate split:
53///
54/// ```text
55///     tensor[hi << n_lo | lo] = eq_lo[lo] * eq_hi[hi]
56/// ```
57///
58/// Multiplication distributes over the sum, so the fold regroups exactly:
59///
60/// ```text
61///     sum_x tensor[x] * row[x]
62///       = sum_hi eq_hi[hi] * ( sum_lo eq_lo[lo] * row[hi, lo] )
63/// ```
64///
65/// Each block of the matrix folds against the low factor alone.
66/// Its result is then scaled once by that block's high-factor entry and merged.
67/// Field addition and multiplication are exact.
68/// So the result is bit-identical to folding against the materialized tensor.
69///
70/// # Why the factored form is faster
71///
72/// Both effects follow from the tables covering only the low factor, so being built once.
73///
74/// Memory traffic:
75///
76/// - The full tensor is never written or read, so only the matrix streams through.
77/// - Folding against a materialized tensor reads one tensor entry per row alongside it.
78/// - That doubles the stream for the same arithmetic.
79///
80/// Lookup count:
81///
82/// - A table built once costs nothing per row, so it can cover eight rows instead of four.
83/// - Eight-row groups halve the lookups and accumulator updates.
84/// - Those updates are what the fold is bound on, not arithmetic or bandwidth.
85/// - A table rebuilt per group cannot widen: 256 entries per group costs more than it saves.
86///
87/// The price is one multiply per row, to scale each block by its high-factor entry.
88///
89/// # Preconditions
90///
91/// * The matrix must have as many rows as the two factors have entries together.
92/// * A block must cover a whole number of packed elements, so the chunking below aligns. The
93///   exception is a matrix that fits inside one packed element, which is a single block.
94/// * The low factor must span at most one row fold chunk, which is 128 rows.
95pub fn fold_1b_rows_for_b128_split<P, Data>(
96	mat: &FieldBuffer<P, Data>,
97	eq_lo: &FieldBuffer<B128>,
98	eq_hi: &FieldBuffer<B128>,
99) -> FieldBuffer<B128>
100where
101	P: PackedField<Scalar = B128>,
102	Data: Deref<Target = [P]>,
103{
104	let log_scalar_bit_width = <B128 as ExtensionField<B1>>::LOG_DEGREE;
105	assert_eq!(mat.log_len(), eq_lo.log_len() + eq_hi.log_len()); // precondition
106	assert!(eq_lo.log_len() >= P::LOG_WIDTH || eq_hi.log_len() == 0); // precondition
107	assert!(eq_lo.log_len() <= LOG_SPLIT_BLOCK); // precondition
108
109	// One subset-sum table per group of eight rows, built once and reused by every hi-block.
110	// A low factor shorter than a full block leaves the trailing weights zero.
111	// So every row past its end is weighted zero, whatever bits that row holds:
112	// rows past the end of a block read as zero, and a matrix narrower than one packed
113	// element leaves lanes that are not rows at all.
114	let lo_tables = RowFoldTables::<B128, N_ROW_TABLES>::new(eq_lo.as_ref());
115
116	// A sub-width low factor still occupies its one packed element, so the block is that element.
117	let block_packed_len = 1 << eq_lo.log_len().saturating_sub(P::LOG_WIDTH);
118
119	(mat.as_ref().par_chunks(block_packed_len), eq_hi.as_ref().par_iter())
120		.into_par_iter()
121		.fold(
122			|| FieldBuffer::zeros(log_scalar_bit_width),
123			|mut acc, (mat_block, &eq_hi_val)| {
124				// Fold this block against the low factor alone:
125				//
126				//     block[col] = sum_lo eq_lo[lo] * bit_col(mat_block[lo])
127				//
128				// A matrix row is 128 single-bit columns, which is one full-width packed row.
129				let mut rows = P::iter_slice(mat_block);
130				let mut sums = ColumnSums::zero();
131
132				// Gather each group's rows out of the packed elements they sit in, lazily, so a
133				// group is folded as soon as it is complete.
134				// Rows past the end of the block stay zero, and lanes past the last row carry
135				// weight zero, so neither contributes.
136				// A field element and a 128-bit row of single-bit scalars share one underlier, so
137				// each view is free.
138				lo_tables.fold_into(
139					iter::repeat_with(|| {
140						array::from_fn(|_| {
141							rows.next()
142								.map(packed_extension::cast_base::<B1, _>)
143								.unwrap_or_default()
144						})
145					})
146					.take(N_ROW_TABLES),
147					&mut sums,
148				);
149
150				// Scale by this block's high-factor entry and merge.
151				// That is 128 multiplies per block, one per row.
152				sums.add_scaled_to(eq_hi_val, acc.as_mut());
153				acc
154			},
155		)
156		// A merge seeded with a partial that already exists never touches a buffer of zeros.
157		// An identity would allocate and zero one accumulator per merge, then add all of it.
158		.reduce_with(|mut lhs, rhs| {
159			for (lhs_i, &rhs_i) in izip!(lhs.as_mut(), rhs.as_ref()) {
160				*lhs_i += rhs_i;
161			}
162			lhs
163		})
164		// An empty matrix yields no partials at all, and folds to zero.
165		.unwrap_or_else(|| FieldBuffer::zeros(log_scalar_bit_width))
166}
167
168/// Builds the ring-switching equality indicator directly from the tensor's two factors.
169///
170/// # Overview
171///
172/// The indicator is the suffix equality tensor, folded bitwise by the row-batching query.
173/// Expanding that tensor and then folding it would read and rewrite a whole `2^n` buffer.
174/// This route produces each entry in one pass instead:
175///
176/// ```text
177///     out[hi << n_lo | lo] = fold(eq_lo[lo] * eq_hi[hi])
178/// ```
179///
180/// The tensor expansion already costs one multiply per entry.
181/// So the fused product changes no operation count.
182/// It removes one full read pass and one full write pass over the `2^n` buffer.
183///
184/// ## Arguments
185///
186/// * `alloc` - the allocator the returned indicator is drawn from
187/// * `eq_lo` - the low tensor factor
188/// * `eq_hi` - the high tensor factor
189/// * `row_batch_query` - the vector every entry is folded bitwise by
190///
191/// ## Preconditions
192///
193/// * `row_batch_query.len()` must equal 128, the extension degree of B128 over B1
194/// * A block must cover a whole number of packed elements, so the chunking below aligns. The
195///   exception is an output that fits inside one packed element, which is a single block.
196pub fn rs_eq_ind_from_factors<A, P>(
197	alloc: &A,
198	eq_lo: &FieldBuffer<B128>,
199	eq_hi: &FieldBuffer<B128>,
200	row_batch_query: &FieldBuffer<B128>,
201) -> FieldVec<P, A>
202where
203	A: Allocator,
204	P: PackedField<Scalar = B128>,
205{
206	assert!(eq_lo.log_len() >= P::LOG_WIDTH || eq_hi.log_len() == 0); // precondition
207	assert_eq!(row_batch_query.log_len(), <B128 as ExtensionField<B1>>::LOG_DEGREE); // precondition
208
209	// The bitwise fold that maps a tensor entry to its indicator entry, shared by every block.
210	let transform = OutputWrappingTransformationFactory::new(
211		InputWrappingTransformationFactory::new(BytewiseLookupTransformationFactory),
212	)
213	.create(row_batch_query.as_ref());
214
215	let log_len = eq_lo.log_len() + eq_hi.log_len();
216	let packed_len = 1usize << log_len.saturating_sub(P::LOG_WIDTH);
217	let block_packed_len = 1 << eq_lo.log_len().saturating_sub(P::LOG_WIDTH);
218
219	// The buffer is written exactly once, block by block, so it starts uninitialized.
220	let mut out = alloc.alloc::<P>(packed_len);
221	(
222		out.spare_capacity_mut()[..packed_len].par_chunks_mut(block_packed_len),
223		eq_hi.as_ref().par_iter(),
224	)
225		.into_par_iter()
226		.for_each(|(out_block, &eq_hi_val)| {
227			// Each slot holds P::WIDTH consecutive entries of this hi-block.
228			// Every entry is one product of the two factors, folded and written once.
229			// A sub-width output leaves one short chunk, whose lanes past the last entry are
230			// filled with zeros rather than left undefined.
231			let lo_chunks = eq_lo.as_ref().chunks(P::WIDTH);
232			for (slot, lo_chunk) in iter::zip(out_block, lo_chunks) {
233				slot.write(P::from_scalars(
234					lo_chunk
235						.iter()
236						.map(|&lo| transform.transform(&(lo * eq_hi_val))),
237				));
238			}
239		});
240	// SAFETY: the block partition covers all `packed_len` slots and every slot was written.
241	unsafe { out.set_len(packed_len) };
242
243	FieldBuffer::new(log_len, out)
244}
245
246/// Expands `point` into the two factors of its equality tensor, low factor first.
247///
248/// The `2^n` tensor itself is never materialized: both consumers read the factors directly.
249/// The low factor spans one row fold chunk, or the whole point when that is shorter.
250fn expand_tensor_factors(point: &[B128]) -> (FieldBuffer<B128>, FieldBuffer<B128>) {
251	let (point_lo, point_hi) = point.split_at(point.len().min(LOG_SPLIT_BLOCK));
252	(eq_ind_partial_eval::<B128>(point_lo), eq_ind_partial_eval::<B128>(point_hi))
253}
254
255/// Output of ring-switching prover.
256pub struct RingSwitchOutput<A: Allocator, P: PackedField> {
257	/// The ring-switching equality indicator MLE (transparent poly for BaseFold).
258	pub rs_eq_ind: FieldVec<P, A>,
259	/// The sumcheck claim.
260	pub sumcheck_claim: P::Scalar,
261}
262
263/// Prove the ring-switching reduction.
264///
265/// Takes the packed witness and evaluation point from shift reduction, and:
266/// 1. Computes partial evaluations s_hat_v
267/// 2. Sends s_hat_v to verifier via channel
268/// 3. Samples row-batching challenges
269/// 4. Computes the ring-switching equality indicator and sumcheck claim
270///
271/// Returns the transparent polynomial and sumcheck claim for BaseFold.
272///
273/// ## Arguments
274///
275/// * `alloc` - the allocator the ring-switching equality indicator is drawn from
276/// * `packed_witness` - the packed witness buffer (B1 polynomial packed into P elements)
277/// * `eval_point` - the evaluation point from shift reduction
278/// * `channel` - the prover channel for sending/sampling
279///
280/// ## Preconditions
281///
282/// * `packed_witness.log_len() + log_packing == eval_point.len()` where log_packing is the base-2
283///   log of the extension degree of B128 over B1 (= 7)
284pub fn prove<A, P, Channel>(
285	alloc: &A,
286	packed_witness: FieldSlice<'_, P>,
287	eval_point: &[B128],
288	channel: &mut Channel,
289) -> RingSwitchOutput<A, P>
290where
291	A: Allocator,
292	P: PackedField<Scalar = B128>,
293	Channel: IPProverChannel<B128>,
294{
295	let log_packing = <B128 as ExtensionField<B1>>::LOG_DEGREE;
296	assert_eq!(packed_witness.log_len() + log_packing, eval_point.len());
297
298	let eval_point_suffix = &eval_point[log_packing..];
299	let (eq_lo, eq_hi) = tracing::debug_span!("Expand evaluation suffix query")
300		.in_scope(|| expand_tensor_factors(eval_point_suffix));
301
302	// Ring-switching partial evaluations (Method of Four Russians)
303	let s_hat_v = tracing::debug_span!("Compute ring-switching partial evaluations")
304		.in_scope(|| fold_1b_rows_for_b128_split(&packed_witness, &eq_lo, &eq_hi));
305	channel.send_many(s_hat_v.as_ref());
306
307	// Basis transpose
308	let s_hat_u = TensorAlgebra::<B1, B128>::new(s_hat_v.as_ref().to_vec())
309		.transpose()
310		.elems;
311
312	// Sample row-batching challenges
313	let r_double_prime = channel.sample_many(log_packing);
314	let eq_r_double_prime = eq_ind_partial_eval::<B128>(&r_double_prime);
315
316	// GF(2^128) reduction is F2-linear, so it commutes with XOR.
317	// Summing 128 wide products then reducing once matches reducing each term first.
318	let sumcheck_claim = inner_product_packed::<B128, B128>(
319		log_packing,
320		s_hat_u.into_iter(),
321		eq_r_double_prime.as_ref().iter().copied(),
322	);
323
324	// Compute ring-switching equality indicator (transparent poly)
325	let rs_eq_ind = tracing::debug_span!("Compute ring-switching equality indicator")
326		.in_scope(|| rs_eq_ind_from_factors::<A, P>(alloc, &eq_lo, &eq_hi, &eq_r_double_prime));
327
328	RingSwitchOutput {
329		rs_eq_ind,
330		sumcheck_claim,
331	}
332}
333
334/// Proves the public segment's evaluation claim.
335///
336/// The shift closes over the public segment as a bit matrix, at `r_j` over the bit within a word
337/// and the low coordinates of `r_y` over the word index. The verifier holds the segment but not
338/// its bits, so the claim is stated here and reduced in two steps:
339///
340/// 1. a ring-switch onto the segment's packed form, leaving the claim `sum_x P(x) A(x)` against the
341///    ring-switching indicator;
342/// 2. a sumcheck over the packed segment's own variables, leaving one evaluation of each factor.
343///
344/// Nothing here is committed, so the verifier finishes on its own: it evaluates the packed
345/// segment's multilinear from the words it holds, and the indicator from its succinct formula.
346///
347/// ## Arguments
348///
349/// * `alloc` - the allocator the packed segment and the indicator are drawn from
350/// * `public_words` - the public segment, unpadded
351/// * `r_j` - the bit-index challenges
352/// * `r_y` - the word-index challenges, of which the segment spans the low ones
353/// * `channel` - the prover channel for sending/sampling
354///
355/// ## Preconditions
356///
357/// * `r_y` must have at least as many coordinates as the packed segment spans words
358pub fn prove_public_eval<A, P, Channel>(
359	alloc: &A,
360	public_words: &[Word],
361	r_j: &[B128],
362	r_y: &[B128],
363	channel: &mut Channel,
364) where
365	A: Allocator,
366	P: PackedField<Scalar = B128>,
367	Channel: IPProverChannel<B128>,
368{
369	// The claim is over the packed segment, so it spans whole field elements: a segment shorter
370	// than one still spans one, reading the words past its end as zero.
371	let log_public_elems = log2_ceil_usize(public_words.len()).saturating_sub(LOG_WORDS_PER_ELEM);
372	let r_y_public = &r_y[..log_public_elems + LOG_WORDS_PER_ELEM];
373
374	channel.send_one(evaluate_words_mle::<B128, B128>(public_words, r_j, r_y_public));
375
376	let packed = pack_witness::<P, _>(alloc, log_public_elems, public_words)
377		.expect("the element count is derived from the words being packed");
378	let RingSwitchOutput {
379		rs_eq_ind,
380		sumcheck_claim,
381	} = prove(alloc, packed.as_view(), &[r_j, r_y_public].concat(), channel);
382
383	// The reduced claim is the sum of the two multilinears' product over the hypercube, which is
384	// what the trace's opening hands to BaseFold. Here it is discharged by the sumcheck alone: the
385	// final evaluations need no message, since the verifier computes both itself.
386	let prover = bivariate_product_prover(alloc, [packed, rs_eq_ind], sumcheck_claim);
387	prove_single(prover, channel);
388}
389
390#[cfg(test)]
391mod test {
392	use binius_compute::GlobalAllocator;
393	use binius_field::{
394		ExtensionField, Field, Ghash128b, PackedField, PackedGhash2x128b, PackedGhash4x128b,
395		PackedSubfield, packed_extension,
396	};
397	use binius_math::{
398		FieldBuffer,
399		inner_product::{inner_product_buffers, inner_product_subfield},
400		multilinear::{eq::eq_ind_partial_eval, evaluate::evaluate_inplace},
401		test_utils::{index_to_hypercube_point, random_field_buffer, random_scalars},
402	};
403	use binius_verifier::{config::B1, ring_switch::eval_rs_eq};
404	use rand::{SeedableRng, rngs::StdRng};
405
406	use super::*;
407
408	type F = Ghash128b;
409
410	// The row fold, written straight from its definition:
411	//
412	//     out[b] = sum_r eq[r] * bit_b(row r)
413	//
414	// Only the `eq.len()` rows the weights address take part.
415	// Lanes past the last row of a sub-width buffer are not rows, so they never appear here.
416	fn naive_fold_1b_rows<P: PackedField<Scalar = F>>(mat: &FieldBuffer<P>, eq: &[F]) -> Vec<F> {
417		let mut out = vec![F::ZERO; <F as ExtensionField<B1>>::DEGREE];
418		for (r, &weight) in eq.iter().enumerate() {
419			let row = mat.get(r);
420			for (bit, out_b) in iter::zip(ExtensionField::<B1>::iter_bases(&row), &mut out) {
421				if bit == B1::ONE {
422					*out_b += weight;
423				}
424			}
425		}
426		out
427	}
428
429	// The split fold must reproduce that definition at every legal split of the point.
430	//
431	//     split:  fold(mat, eq(lo), eq(hi)) with the tensor never materialized
432	//
433	// The matrix is random packed elements, so a sub-width buffer holds unrelated data in the
434	// lanes past its last row. The fold must weight those zero and reach the same result.
435	fn check_split_fold_matches_definition<P: PackedField<Scalar = F>>(log_len: usize, seed: u64) {
436		let mut rng = StdRng::seed_from_u64(seed);
437		let mat = random_field_buffer::<P>(&mut rng, log_len);
438		let point: Vec<F> = random_scalars(&mut rng, log_len);
439
440		let expected = naive_fold_1b_rows(&mat, eq_ind_partial_eval::<F>(&point).as_ref());
441
442		// The low factor must cover whole packed elements, unless the matrix is one element.
443		// The ceiling is the 128-row chunk the row fold consumes.
444		for split_at in P::LOG_WIDTH.min(log_len)..=log_len.min(LOG_SPLIT_BLOCK) {
445			let (point_lo, point_hi) = point.split_at(split_at);
446			let eq_lo = eq_ind_partial_eval::<F>(point_lo);
447			let eq_hi = eq_ind_partial_eval::<F>(point_hi);
448
449			let split = fold_1b_rows_for_b128_split(&mat, &eq_lo, &eq_hi);
450			assert_eq!(
451				split.as_ref(),
452				expected.as_slice(),
453				"log_len={log_len}, split_at={split_at}"
454			);
455		}
456	}
457
458	#[test]
459	fn test_split_fold_matches_definition() {
460		// Row counts below, at and above one packed element, for each packing width.
461		for (i, log_len) in [0, 1, 2, 6, 7, 8].into_iter().enumerate() {
462			let seed = i as u64;
463			check_split_fold_matches_definition::<F>(log_len, seed);
464			check_split_fold_matches_definition::<PackedGhash2x128b>(log_len, seed);
465			check_split_fold_matches_definition::<PackedGhash4x128b>(log_len, seed);
466		}
467	}
468
469	// The indicator must match the verifier's succinct formula for the same polynomial.
470	//
471	//     out[i] = A(point, i)
472	//
473	// `eval_rs_eq` evaluates `A` through the tensor algebra, independently of how the prover
474	// builds the buffer, so it pins the fused route at every legal split of the point.
475	fn check_rs_eq_ind_from_factors<P: PackedField<Scalar = F>>(log_len: usize, seed: u64) {
476		let mut rng = StdRng::seed_from_u64(seed);
477		let point: Vec<F> = random_scalars(&mut rng, log_len);
478		let row_batching_challenges: Vec<F> =
479			random_scalars(&mut rng, <F as ExtensionField<B1>>::LOG_DEGREE);
480		let row_batch_query = eq_ind_partial_eval::<F>(&row_batching_challenges);
481
482		// The indicator has no chunk ceiling, so the low factor may span the whole point.
483		for split_at in P::LOG_WIDTH.min(log_len)..=log_len {
484			let (point_lo, point_hi) = point.split_at(split_at);
485			let eq_lo = eq_ind_partial_eval::<F>(point_lo);
486			let eq_hi = eq_ind_partial_eval::<F>(point_hi);
487
488			let rs_eq_ind =
489				rs_eq_ind_from_factors::<_, P>(&GlobalAllocator, &eq_lo, &eq_hi, &row_batch_query);
490
491			for index in 0..1 << log_len {
492				let expected = eval_rs_eq::<F>(
493					&point,
494					&index_to_hypercube_point::<F>(log_len, index),
495					row_batch_query.as_ref(),
496				);
497				assert_eq!(rs_eq_ind.get(index), expected, "split_at={split_at}, index={index}");
498			}
499
500			// A sub-width buffer's trailing lanes are not entries, and the sumcheck provers
501			// downstream read them as zero, so the write must leave them zero.
502			let trailing = rs_eq_ind.as_ref()[0].iter().skip(1 << log_len);
503			assert!(trailing.take(P::WIDTH).all(|lane| lane == F::ZERO), "split_at={split_at}");
504		}
505	}
506
507	#[test]
508	fn test_rs_eq_ind_from_factors() {
509		// Entry counts below, at and above one packed element, for each packing width.
510		for (i, log_len) in [0, 1, 2, 6, 7, 8].into_iter().enumerate() {
511			let seed = i as u64;
512			check_rs_eq_ind_from_factors::<F>(log_len, seed);
513			check_rs_eq_ind_from_factors::<PackedGhash2x128b>(log_len, seed);
514			check_rs_eq_ind_from_factors::<PackedGhash4x128b>(log_len, seed);
515		}
516	}
517
518	#[test]
519	fn test_out_of_range_evaluation() {
520		let mut rng = StdRng::from_seed([0; 32]);
521
522		// `eval_rs_eq` must agree with the indicator off the hypercube as well:
523		//
524		//     eval_rs_eq(z, x, q) = sum_i eq(x)[i] * rs_eq_ind[i]
525		let n_vars_big_field = 3;
526
527		// setup ring switch eq mle
528		let z_vals: Vec<F> = random_scalars(&mut rng, n_vars_big_field);
529
530		let row_batching_challenges: Vec<F> =
531			random_scalars(&mut rng, <F as ExtensionField<B1>>::LOG_DEGREE);
532
533		let row_batching_expanded_query: FieldBuffer<F> =
534			eq_ind_partial_eval(&row_batching_challenges);
535
536		let (eq_lo, eq_hi) = expand_tensor_factors(&z_vals);
537		let rs_eq = rs_eq_ind_from_factors::<_, F>(
538			&GlobalAllocator,
539			&eq_lo,
540			&eq_hi,
541			&row_batching_expanded_query,
542		);
543
544		// out of range eval point
545		let eval_point: Vec<F> = random_scalars(&mut rng, n_vars_big_field);
546
547		// compare eval against inner product w/ eq ind mle of eval point
548
549		let tensor_expanded_eval_point = eq_ind_partial_eval::<F>(&eval_point);
550		let expected_eval = inner_product_buffers(&rs_eq, &tensor_expanded_eval_point);
551
552		let actual_eval =
553			eval_rs_eq::<F>(&z_vals, &eval_point, row_batching_expanded_query.as_ref());
554
555		assert_eq!(expected_eval, actual_eval);
556	}
557
558	#[test]
559	fn test_row_fold_composes_into_the_claim() {
560		let mut rng = StdRng::seed_from_u64(0);
561
562		type P = PackedGhash2x128b;
563
564		// The prover's row fold and the verifier's partial evaluation compose into the claim:
565		//
566		//     <eq(point), bits>  ==  evaluate(s_hat_v, point[..log_degree])
567		//
568		// where `s_hat_v[b] = sum_r eq(point[log_degree..])[r] * bit_b(row r)`.
569		// The low coordinates of the point index the bits within a row, the high ones the rows.
570		let n = 10;
571		let log_degree = <F as ExtensionField<B1>>::LOG_DEGREE;
572
573		// A random B1 matrix of 2^(n + log_degree) bits, and a point over the same index space.
574		let bit_matrix = random_field_buffer::<PackedSubfield<P, B1>>(&mut rng, n + log_degree);
575		let eval_point: Vec<F> = random_scalars(&mut rng, n + log_degree);
576		let (prefix, suffix) = eval_point.split_at(log_degree);
577
578		// Reference: expand the whole point and contract it against every bit.
579		let full_tensor = eq_ind_partial_eval::<F>(&eval_point);
580		let expected = inner_product_subfield(
581			PackedField::iter_slice(bit_matrix.as_ref()),
582			PackedField::iter_slice(full_tensor.as_ref()),
583		);
584
585		// Prover route: fold the rows against the suffix factors, then evaluate at the prefix.
586		// Viewing 128 single-bit scalars as one field element is a reinterpretation, so the rows
587		// are the same memory.
588		let mat = FieldBuffer::<P>::new(
589			n,
590			bit_matrix
591				.as_ref()
592				.iter()
593				.map(|&bits_packed| packed_extension::cast_ext::<B1, P>(bits_packed))
594				.collect(),
595		);
596		let (eq_lo, eq_hi) = expand_tensor_factors(suffix);
597		let s_hat_v = fold_1b_rows_for_b128_split(&mat, &eq_lo, &eq_hi);
598
599		assert_eq!(evaluate_inplace(s_hat_v, prefix), expected);
600	}
601}