Skip to main content

binius_prover/protocols/shift/
shift_ind.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! The sumcheck rounds binding the bit index the shift indicators read, and the shift indicator
5//! partial evaluations they run over.
6
7use binius_compute::Allocator;
8use binius_core::{ShiftVariant, word::Word};
9use binius_field::{BinaryField, FieldOps, PackedField};
10use binius_ip_prover::{
11	channel::IPProverChannel,
12	sumcheck::{ProveSingleOutput, bivariate_product_prover, prove_single},
13};
14use binius_math::{
15	FieldBuffer, FieldVec,
16	inner_product::inner_product,
17	multilinear::eq::{eq_ind_partial_eval, eq_ind_partial_eval_scalars},
18};
19use binius_verifier::protocols::shift::LOG_SHIFT_VARIANT_COUNT;
20
21/// The number of bit variables a half-word (`*32`) shift variant acts over.
22const HALF_WORD_LOG_BITS: usize = Word::LOG_BITS - 1;
23
24/// One shift's amount and variant challenges.
25#[derive(Debug, Clone)]
26pub struct ShiftChallenge<F> {
27	/// The shift amount.
28	pub(crate) amount: Vec<F>,
29	/// The shift variant.
30	pub(crate) variant: Vec<F>,
31}
32
33impl<F> ShiftChallenge<F> {
34	/// Builds a shift challenge from its two axes.
35	pub const fn new(amount: Vec<F>, variant: Vec<F>) -> Self {
36		debug_assert!(amount.len() == Word::LOG_BITS, "one challenge per bit position of a word");
37		debug_assert!(variant.len() == LOG_SHIFT_VARIANT_COUNT, "one challenge per shift variant");
38		Self { amount, variant }
39	}
40}
41
42/// The point a shift indicator is read at.
43#[derive(Debug, Clone, Copy)]
44pub struct ShiftChallengePoint<'a, F> {
45	/// Which bit position of a word the point selects.
46	bit: &'a [F],
47	/// Which shift, amount and variant together, the point selects.
48	shift: &'a ShiftChallenge<F>,
49}
50
51impl<'a, F: BinaryField> ShiftChallengePoint<'a, F> {
52	/// Builds a challenge point from a bit position and a shift.
53	pub const fn new(bit: &'a [F], shift: &'a ShiftChallenge<F>) -> Self {
54		debug_assert!(bit.len() == Word::LOG_BITS, "one challenge per bit position of a word");
55		Self { bit, shift }
56	}
57
58	/// Builds the shift indicators over the bit index at this point, one scalar per bit index.
59	///
60	/// The variant axis is folded into a single multilinear here.
61	fn indicator(&self) -> Vec<F> {
62		let bit = self.bit;
63		let amount = self.shift.amount.as_slice();
64		let variant = self.shift.variant.as_slice();
65
66		let (sigma, sigma_prime) = partial_eval_sigmas(bit, amount);
67		let sigma_transpose = partial_eval_sigmas_transpose(bit, amount);
68		let phi = partial_eval_phi(amount);
69		// The equality indicator selecting the sign position, the top input bit.
70		let sign_position: F = bit.iter().copied().product();
71
72		// A half-word variant applies the same four rules to each 32-bit half, reading only the
73		// low bits of the shift amount. The two halves never mix, so each indicator also carries
74		// the equality of the halves the two bit indices fall in.
75		let (sigma32, sigma32_prime) =
76			partial_eval_sigmas(&bit[..HALF_WORD_LOG_BITS], &amount[..HALF_WORD_LOG_BITS]);
77		let sigma32_transpose = partial_eval_sigmas_transpose(
78			&bit[..HALF_WORD_LOG_BITS],
79			&amount[..HALF_WORD_LOG_BITS],
80		);
81		let phi32 = partial_eval_phi(&amount[..HALF_WORD_LOG_BITS]);
82		let sign_position32: F = bit[..HALF_WORD_LOG_BITS].iter().copied().product();
83		let same_half = eq_ind_partial_eval::<F>(&bit[HALF_WORD_LOG_BITS..]);
84
85		let variant_tensor = eq_ind_partial_eval::<F>(variant);
86		(0..Word::BITS)
87			.map(|index| {
88				let (half, low) = (index >> HALF_WORD_LOG_BITS, index % (1 << HALF_WORD_LOG_BITS));
89				let same_half = same_half.as_ref()[half];
90				// The eight indicators at this bit index, one per shift variant.
91				let shift_inds = ShiftVariant::ALL.map(|shift_variant| match shift_variant {
92					ShiftVariant::Sll => sigma_transpose[index],
93					ShiftVariant::Slr => sigma[index],
94					ShiftVariant::Sar => sigma[index] + sign_position * phi[index],
95					ShiftVariant::Rotr => sigma[index] + sigma_prime[index],
96					ShiftVariant::Sll32 => same_half * sigma32_transpose[low],
97					ShiftVariant::Srl32 => same_half * sigma32[low],
98					ShiftVariant::Sra32 => {
99						same_half * (sigma32[low] + sign_position32 * phi32[low])
100					}
101					ShiftVariant::Rotr32 => same_half * (sigma32[low] + sigma32_prime[low]),
102				});
103				inner_product(shift_inds, variant_tensor.as_ref().iter().copied())
104			})
105			.collect()
106	}
107}
108
109/// Phase 3 of the shift reduction's sumcheck: the [`Word::LOG_BITS`] rounds binding the bit
110/// index the shift indicators read.
111///
112/// Phases 1 and 2 leave the claim
113///
114/// $$
115/// \beta = h(r_j, r_s, r_v) \cdot G, \qquad
116/// h(r_j, r_s, r_v) = \sum_i \widetilde{L}(i) \cdot
117///     \sum_{\text{op}} \widetilde{eq}(r_v, \text{op}) \cdot \text{ind}_{\text{op}}(i, r_j, r_s),
118/// $$
119///
120/// with $G = g(r_j, r_s, r_v)$ the sum over the word index that phase 4 goes on to bind. Unrolling
121/// $h$ exposes these rounds as a sumcheck over two multilinears in the bit index — a weight vector
122/// and the interpolated shift indicators — with $G$ riding along as a constant.
123///
124/// The two are held apart rather than multiplied together, which is what keeps the round
125/// polynomials degree 2. The constant is folded into the weights, so the pair sums to $\beta$ and
126/// the rounds are the ones the verifier's single sumcheck expects. Phase 4 then scales its
127/// monster multilinear by the product of the two evaluations these rounds reduce their factors to,
128/// which [`ShiftIndOutput`] reports separately.
129///
130/// The weights and the point the indicator is read at are the caller's, not this type's: a
131/// reduction peeling two shifts runs these rounds once per shift slot, differing only in those two
132/// arguments.
133pub struct ShiftIndSumcheck<P: PackedField, A: Allocator> {
134	/// The weights over the bit index, scaled by `G`.
135	scaled_weights: FieldVec<P, A>,
136	/// The shift indicators interpolated over the shift variant, over the bit index.
137	shift_ind: FieldVec<P, A>,
138	/// The unscaled weights, kept to evaluate them at the challenge point.
139	weights: Vec<P::Scalar>,
140	/// The claim these rounds start from: `h(point) * G`.
141	beta: P::Scalar,
142}
143
144/// What one run of these rounds leaves the phases after it.
145///
146/// The two factor evaluations are reported apart rather than as their product. A reduction peeling
147/// two shifts runs these rounds once per slot and needs them separately: the indicator evaluation
148/// alone is the constant the next run carries, while the products of both runs are what scale the
149/// monster multilinear at the end.
150#[derive(Debug, Clone)]
151pub struct ShiftIndOutput<F> {
152	/// The weight vector at the challenge point, `weights(r_i)`.
153	pub weights_eval: F,
154	/// The interpolated shift indicator at the challenge point,
155	/// `shift_ind(r_i, point)`. The verifier recomputes it from the point alone.
156	pub ind_eval: F,
157	/// The claim the next phase proves: the two evaluations above times the carried constant.
158	pub eval: F,
159	/// The point these rounds bound, in evaluation order.
160	pub point: Vec<F>,
161}
162
163impl<F: BinaryField, P: PackedField<Scalar = F>, A: Allocator> ShiftIndSumcheck<P, A> {
164	/// Builds the two multilinears the rounds run over, from the weights the caller holds and
165	/// phase 1's challenges.
166	///
167	/// # Arguments
168	///
169	/// - `weights`: the weight vector over the bit index, one entry per bit of a word. The
170	///   reduction supplies the oblong Lagrange evaluations at the univariate challenge.
171	/// - `point`: the point the shift indicator is read at — phase 1's challenges for the input bit
172	///   position, the shift amount and the shift variant.
173	/// - `g_eval`: `g(point)`, the constant these rounds carry.
174	///
175	/// # Panics
176	///
177	/// Panics unless the weights hold one entry per bit position of a word.
178	pub fn new(alloc: &A, weights: &[F], point: &ShiftChallengePoint<'_, F>, g_eval: F) -> Self {
179		assert_eq!(weights.len(), Word::BITS, "the weights are indexed by bit position");
180
181		let shift_ind = point.indicator();
182
183		// Folding the constant into one of the two factors makes the pair sum to the incoming
184		// claim, so the standard bivariate-product prover emits the right rounds.
185		let scaled_weights = weights
186			.iter()
187			.map(|&weight| weight * g_eval)
188			.collect::<Vec<_>>();
189		let beta = inner_product(scaled_weights.iter().copied(), shift_ind.iter().copied());
190
191		Self {
192			scaled_weights: FieldBuffer::from_values_in(alloc, &scaled_weights),
193			shift_ind: FieldBuffer::from_values_in(alloc, &shift_ind),
194			weights: weights.to_vec(),
195			beta,
196		}
197	}
198
199	/// The claim these rounds start from, which phase 2 reduced to.
200	pub const fn beta(&self) -> F {
201		self.beta
202	}
203
204	/// Proves the [`Word::LOG_BITS`] rounds binding the bit index.
205	pub fn prove(self, channel: &mut impl IPProverChannel<F>, alloc: &A) -> ShiftIndOutput<F> {
206		let Self {
207			scaled_weights,
208			shift_ind,
209			weights,
210			beta,
211		} = self;
212
213		let prover = bivariate_product_prover(alloc, [scaled_weights, shift_ind], beta);
214		let ProveSingleOutput {
215			multilinear_evals,
216			mut challenges,
217		} = prove_single(prover, channel);
218		challenges.reverse();
219
220		let [scaled_weights_eval, shift_ind_eval] = multilinear_evals
221			.try_into()
222			.expect("prover has 2 multilinear polynomials");
223
224		// The carried constant rides in the first evaluation, so the unscaled weights are evaluated
225		// at the challenge point separately — the same 64-term inner product the verifier runs.
226		let weights_eval = inner_product(weights, eq_ind_partial_eval_scalars(&challenges));
227
228		ShiftIndOutput {
229			weights_eval,
230			ind_eval: shift_ind_eval,
231			eval: scaled_weights_eval * shift_ind_eval,
232			point: challenges,
233		}
234	}
235}
236
237/// Partial evaluation of the shift indicator helper polynomials $\sigma, \sigma'$ over all i on the
238/// hypercube.
239///
240/// Given fixed j and s, computes sigma and sigma_prime for all possible i values.
241/// Returns (sigma, sigma_prime) as Vecs of length `1 << bit.len()`.
242fn partial_eval_sigmas<E: FieldOps>(bit: &[E], amount: &[E]) -> (Vec<E>, Vec<E>) {
243	assert_eq!(bit.len(), amount.len(), "the two axes must have the same length");
244
245	let n = bit.len();
246	let mut sigma = vec![E::zero(); 1 << n];
247	let mut sigma_prime = vec![E::zero(); 1 << n];
248	sigma[0] = E::one();
249
250	// Process each bit position
251	for k in 0..n {
252		let j_k = bit[k].clone();
253		let s_k = amount[k].clone();
254
255		// Precompute boolean combinations for this bit
256		let both = j_k.clone() * &s_k;
257		let j_one_s = j_k.clone() - &both; // j_k * (1 - s_k)
258		let one_j_s = s_k.clone() - &both; // (1 - j_k) * s_k
259		let xor = j_k + s_k;
260		let eq = E::one() + &xor;
261
262		// Update arrays for this bit position
263		for i in 0..(1 << k) {
264			// Update upper halves first (i_k = 1)
265			sigma[(1 << k) | i] = j_one_s.clone() * &sigma[i];
266			sigma_prime[(1 << k) | i] = one_j_s.clone() * &sigma[i] + eq.clone() * &sigma_prime[i];
267
268			// Update lower halves (i_k = 0)
269			let sigma_i = sigma[i].clone();
270			let sigma_prime_i = sigma_prime[i].clone();
271			sigma[i] = eq.clone() * &sigma_i + j_one_s.clone() * &sigma_prime_i;
272			sigma_prime[i] = sigma_prime_i * &one_j_s;
273		}
274	}
275
276	(sigma, sigma_prime)
277}
278
279/// Partial evaluation of the shift indicator helper polynomial $\phi$ over all i on the hypercube.
280///
281/// Given fixed s, computes phi for all possible i values.
282fn partial_eval_phi<E: FieldOps>(amount: &[E]) -> Vec<E> {
283	let n = amount.len();
284	let mut phi = vec![E::zero(); 1 << n];
285
286	// Process each bit position
287	for k in 0..n {
288		let s_k = amount[k].clone();
289
290		// Update arrays for this bit position
291		for i in 0..(1 << k) {
292			// Update for i_k = 1
293			phi[(1 << k) | i] = s_k.clone() + (E::one() + &s_k) * &phi[i];
294			let temp = phi[(1 << k) | i].clone() - &s_k;
295			phi[i] += &temp;
296		}
297	}
298
299	phi
300}
301
302/// Partial evaluation of transposed sigma for SLL.
303///
304/// Since sll_ind(i, j, s) = srl_ind(j, i, s), this computes sigma with i and j swapped.
305fn partial_eval_sigmas_transpose<E: FieldOps>(bit: &[E], amount: &[E]) -> Vec<E> {
306	assert_eq!(bit.len(), amount.len(), "the two axes must have the same length");
307
308	let n = bit.len();
309	let mut sigma_transpose = vec![E::zero(); 1 << n];
310	let mut sigma_transpose_prime = vec![E::zero(); 1 << n];
311	sigma_transpose[0] = E::one();
312
313	// Process each bit position
314	for k in 0..n {
315		let j_k = bit[k].clone();
316		let s_k = amount[k].clone();
317
318		// Precompute boolean combinations for this bit (with i and j swapped)
319		let both = j_k.clone() * &s_k;
320		let xor = j_k + s_k;
321		let eq = E::one() + &xor;
322		let zero = eq.clone() + &both;
323
324		// Update arrays for this bit position
325		for i in 0..(1 << k) {
326			// Update for i_k = 1
327			sigma_transpose[(1 << k) | i] =
328				xor.clone() * &sigma_transpose[i] + zero.clone() * &sigma_transpose_prime[i];
329			sigma_transpose_prime[(1 << k) | i] = both.clone() * &sigma_transpose_prime[i];
330
331			// Update for i_k = 0
332			let sigma_t = sigma_transpose[i].clone();
333			sigma_transpose_prime[i] =
334				both.clone() * &sigma_t + xor.clone() * &sigma_transpose_prime[i];
335			sigma_transpose[i] = zero.clone() * &sigma_t;
336		}
337	}
338
339	sigma_transpose
340}
341
342#[cfg(test)]
343mod tests {
344	use std::array;
345
346	use binius_field::{Field, Ghash128b as B128};
347	use binius_math::{
348		BinarySubspace, multilinear::eq::eq_ind_partial_eval_scalars, test_utils::random_scalars,
349		univariate::EvaluationDomain,
350	};
351	use binius_verifier::protocols::shift::evaluate_shift_inds;
352	use rand::{SeedableRng, rngs::StdRng};
353
354	use super::*;
355
356	// Ground truth for a shift-indicator MLE, independent of the recurrence under test.
357	//
358	// Fix j, s to the two challenge vectors passed in.
359	//
360	// Over the hypercube in i, the indicator's MLE expands over the (j, s) cube as:
361	//     mle[i] = sum_{j, s in {0,1}^n : cond(i, j, s)} eq(bit, j) * eq(amount, s)
362	fn reference_indicator(
363		bit: &[B128],
364		amount: &[B128],
365		cond: impl Fn(usize, usize, usize) -> bool,
366	) -> Vec<B128> {
367		let n = bit.len();
368		// eq_bit[j] = eq(bit, j), eq_amount[s] = eq(amount, s).
369		//
370		// Both index little-endian, matching the recurrence's bit order.
371		let eq_bit = eq_ind_partial_eval_scalars(bit);
372		let eq_amount = eq_ind_partial_eval_scalars(amount);
373
374		(0..1 << n)
375			.map(|i| {
376				let mut acc = B128::ZERO;
377				for j in 0..1 << n {
378					for s in 0..1 << n {
379						if cond(i, j, s) {
380							acc += eq_bit[j] * eq_amount[s];
381						}
382					}
383				}
384				acc
385			})
386			.collect()
387	}
388
389	// Draw a pseudo-random challenge (r_j, r_s).
390	// The fixed seed keeps failures reproducible.
391	fn challenges(n: usize) -> (Vec<B128>, Vec<B128>) {
392		let mut rng = StdRng::seed_from_u64(0);
393		(random_scalars(&mut rng, n), random_scalars(&mut rng, n))
394	}
395
396	#[test]
397	fn srl_matches_reference() {
398		// srl: output bit i reads input bit j = i + s.
399		// Bits shifted past the top vanish, since no such j is in range.
400		let (bit, amount) = challenges(6);
401		let (sigma, _) = partial_eval_sigmas(&bit, &amount);
402		assert_eq!(sigma, reference_indicator(&bit, &amount, |i, j, s| j == i + s));
403	}
404
405	#[test]
406	fn sll_matches_reference() {
407		// sll is the transpose of srl.
408		// Output bit i = j + s reads input bit j.
409		let (bit, amount) = challenges(6);
410		let sigma_transpose = partial_eval_sigmas_transpose(&bit, &amount);
411		assert_eq!(sigma_transpose, reference_indicator(&bit, &amount, |i, j, s| i == j + s));
412	}
413
414	#[test]
415	fn sra_matches_reference() {
416		// sra behaves like srl within range.
417		// Past the shift, the sign bit j = 2^n - 1 fills every position.
418		let (bit, amount) = challenges(6);
419		let n = bit.len();
420		let (sigma, _) = partial_eval_sigmas(&bit, &amount);
421		let phi = partial_eval_phi(&amount);
422		// The product of every bit challenge is the eq-indicator selecting the all-ones sign
423		// position j = 2^n - 1.
424		let j_product: B128 = bit.iter().copied().product();
425		let sra: Vec<_> = (0..1 << n).map(|i| sigma[i] + j_product * phi[i]).collect();
426		assert_eq!(
427			sra,
428			reference_indicator(&bit, &amount, |i, j, s| j == (i + s).min((1 << n) - 1))
429		);
430	}
431
432	#[test]
433	fn rotr_matches_reference() {
434		// rotr wraps bits leaving the bottom back to the top.
435		// So j = (i + s) mod 2^n.
436		let (bit, amount) = challenges(6);
437		let n = bit.len();
438		let (sigma, sigma_prime) = partial_eval_sigmas(&bit, &amount);
439		let rotr: Vec<_> = (0..1 << n).map(|i| sigma[i] + sigma_prime[i]).collect();
440		assert_eq!(rotr, reference_indicator(&bit, &amount, |i, j, s| j == (i + s) % (1 << n)));
441	}
442
443	/// The multilinear the prover sums over and the point evaluation the verifier checks it
444	/// against must be the same polynomial: over the bit-index hypercube, the two agree.
445	#[test]
446	fn build_matches_the_verifier_point_evaluation() {
447		let mut rng = StdRng::seed_from_u64(0);
448		let bit = random_scalars::<B128>(&mut rng, Word::LOG_BITS);
449		let amount = random_scalars::<B128>(&mut rng, Word::LOG_BITS);
450		let variant = random_scalars::<B128>(&mut rng, LOG_SHIFT_VARIANT_COUNT);
451
452		let shift = ShiftChallenge::new(amount.clone(), variant.clone());
453		let point = ShiftChallengePoint::new(&bit, &shift);
454		let shift_ind = point.indicator();
455		let variant_tensor = eq_ind_partial_eval_scalars(&variant);
456
457		for (index, &value) in shift_ind.iter().enumerate() {
458			// The bit-index hypercube vertex at `index`.
459			let r_i: [B128; Word::LOG_BITS] = array::from_fn(|bit_index| {
460				if (index >> bit_index) & 1 == 1 {
461					B128::ONE
462				} else {
463					B128::ZERO
464				}
465			});
466			let expected = inner_product(
467				evaluate_shift_inds(&r_i, &bit, &amount),
468				variant_tensor.iter().copied(),
469			);
470			assert_eq!(value, expected, "bit index {index}");
471		}
472	}
473
474	/// The claim phase 3 starts from is the weights contracted against the indicator multilinear,
475	/// scaled by the constant it carries.
476	#[test]
477	fn claimed_sum_is_the_weighted_indicator_sum() {
478		use binius_compute::GlobalAllocator;
479		use binius_field::{PackedGhash2x128b, Random, Rijndael8b};
480
481		type P = PackedGhash2x128b;
482
483		let mut rng = StdRng::seed_from_u64(1);
484		let r_zhat_prime = B128::random(&mut rng);
485		let bit = random_scalars::<B128>(&mut rng, Word::LOG_BITS);
486		let amount = random_scalars::<B128>(&mut rng, Word::LOG_BITS);
487		let variant = random_scalars::<B128>(&mut rng, LOG_SHIFT_VARIANT_COUNT);
488
489		let g_eval = B128::random(&mut rng);
490		let subspace = BinarySubspace::<Rijndael8b>::with_dim(Word::LOG_BITS).isomorphic::<B128>();
491		let l_tilde = subspace.lagrange_evals(&r_zhat_prime);
492		let shift = ShiftChallenge::new(amount, variant);
493		let point = ShiftChallengePoint::new(&bit, &shift);
494		let sumcheck = ShiftIndSumcheck::<P, _>::new(&GlobalAllocator, &l_tilde, &point, g_eval);
495
496		let expected = g_eval * inner_product(l_tilde, point.indicator());
497		assert_eq!(sumcheck.beta(), expected);
498	}
499}