Skip to main content

binius_prover/protocols/shift/
monster.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::iter;
5
6use binius_compute::{Allocator, VecLike};
7use binius_core::{ShiftVariant, constraint_system::Shift, word::Word};
8use binius_field::{BinaryField, Field, PackedField};
9use binius_math::{FieldBuffer, FieldVec, multilinear::eq::eq_ind_partial_eval};
10use tracing::instrument;
11
12use super::{phase_1::SHIFT_OPERATOR_LOG_LEN, shift_ind::ShiftChallenge};
13
14/// The width the half-word (`*32`) shift variants act over.
15const HALF_WORD_BITS: usize = 32;
16
17/// The equality-indicator weights of a shift sequence's outer slot.
18///
19/// The sequence weight factorizes across its two slots, so each slot carries its own two tensors:
20/// one over the variant axis, one over the amount axis.
21/// Keeping them apart holds the weights at `2 * SHIFT_COUNT` entries rather than `SHIFT_COUNT^2`.
22pub(super) struct OuterSlotWeights<F: Field> {
23	/// The equality indicator over the outer variant axis, one weight per shift variant.
24	variant: FieldBuffer<F>,
25	/// The equality indicator over the outer amount axis, one weight per shift amount.
26	amount: FieldBuffer<F>,
27}
28
29impl<F: Field> OuterSlotWeights<F> {
30	/// The equality indicators of the outer slot's challenge point.
31	pub(super) fn new(outer: &ShiftChallenge<F>) -> Self {
32		Self {
33			variant: eq_ind_partial_eval::<F>(&outer.variant),
34			amount: eq_ind_partial_eval::<F>(&outer.amount),
35		}
36	}
37
38	/// The weight one outer shift contributes to its sequence's scalar.
39	#[inline]
40	pub(super) fn weight(&self, shift: Shift) -> F {
41		self.variant.as_ref()[shift.variant as usize] * self.amount.as_ref()[shift.amount as usize]
42	}
43}
44
45/// Writes the row of a shift operator table that one `(variant, amount)` pair contributes.
46///
47/// The row holds, for each bit position, the weight that pair moves there:
48///
49/// ```text
50///     row[j] = sum_k psi(k) * shift-ind_variant(k, j, amount)
51/// ```
52///
53/// This is the one place that says what a variant does to a weight vector:
54/// - Logical left and logical right move the weights and leave zeros behind.
55/// - Arithmetic right piles every weight that falls off the end onto the sign position.
56/// - Rotate wraps them around instead.
57/// - The half-word forms apply the same rule to each 32-bit half, reading only the low 5 bits of
58///   the amount.
59///
60/// # Arguments
61///
62/// The amount is an index over the reduction's amount axis rather than a validated
63/// [`Shift`] amount: that axis spans `Word::BITS` for every
64/// variant, and a half-word variant reads it modulo its own 32-bit width.
65///
66/// Every cell of `row` is written, so the caller need not zero it first. A caller reading one
67/// slice at a time can therefore carry a single scratch row across every pair it visits.
68///
69/// # Panics
70///
71/// Panics unless the row and the weights each hold one entry per bit position of a word.
72pub fn shift_operator_row<F: Field>(
73	variant: ShiftVariant,
74	amount: usize,
75	row: &mut [F],
76	psi: &[F],
77) {
78	assert_eq!(row.len(), Word::BITS, "the row is indexed by bit position");
79	assert_eq!(psi.len(), Word::BITS, "the weights are indexed by bit position");
80
81	// A half-word variant repeats the full-width rule over each half, so both share one closure.
82	let halves = |row: &mut [F], rule: fn(usize, &mut [F], &[F])| {
83		let amount = amount % HALF_WORD_BITS;
84		for (row_half, psi_half) in
85			iter::zip(row.chunks_mut(HALF_WORD_BITS), psi.chunks(HALF_WORD_BITS))
86		{
87			rule(amount, row_half, psi_half);
88		}
89	};
90
91	fn sll<F: Field>(amount: usize, row: &mut [F], psi: &[F]) {
92		let width = row.len();
93		row[..width - amount].copy_from_slice(&psi[amount..]);
94		// The positions the weights vacate take no weight at all.
95		row[width - amount..].fill(F::ZERO);
96	}
97	fn srl<F: Field>(amount: usize, row: &mut [F], psi: &[F]) {
98		let width = row.len();
99		row[..amount].fill(F::ZERO);
100		row[amount..].copy_from_slice(&psi[..width - amount]);
101	}
102	fn sar<F: Field>(amount: usize, row: &mut [F], psi: &[F]) {
103		let width = row.len();
104		srl(amount, row, psi);
105		// Every position past the shift reads the sign bit, so their weights pile onto it.
106		row[width - 1] += psi[width - amount..].iter().sum::<F>();
107	}
108	fn rotr<F: Field>(amount: usize, row: &mut [F], psi: &[F]) {
109		let width = row.len();
110		row[..amount].copy_from_slice(&psi[width - amount..]);
111		row[amount..].copy_from_slice(&psi[..width - amount]);
112	}
113
114	match variant {
115		ShiftVariant::Sll => sll(amount, row, psi),
116		ShiftVariant::Slr => srl(amount, row, psi),
117		ShiftVariant::Sar => sar(amount, row, psi),
118		ShiftVariant::Rotr => rotr(amount, row, psi),
119		ShiftVariant::Sll32 => halves(row, sll),
120		ShiftVariant::Srl32 => halves(row, srl),
121		ShiftVariant::Sra32 => halves(row, sar),
122		ShiftVariant::Rotr32 => halves(row, rotr),
123	}
124}
125
126/// Pushes one weight vector through every shift.
127///
128/// A shift indicator says whether an output bit reads a given input bit at a given amount.
129/// This contracts it on the output index, against the supplied weights:
130///
131/// ```text
132///     T[psi](j, s, o) = sum_k psi(k) * shift-ind_op(o)(k, j, s)
133/// ```
134///
135/// At most one input bit feeds each output bit.
136/// So a slice at fixed `(s, o)` is the weights moved by that shift, not a matrix applied to them.
137///
138/// The reduction contracts the indicator once per slot of a shift sequence.
139/// Both contractions are this operator:
140///
141/// - the first carries the oblong weights to the bits of the intermediate word;
142/// - the second carries that result down to the witness bit.
143///
144/// # Returns
145///
146/// One multilinear over [`SHIFT_OPERATOR_LOG_LEN`] variables, indexed from the low variables up:
147///
148/// ```text
149///     low     Word::LOG_BITS             the bit position
150///     middle  Word::LOG_BITS             the shift amount
151///     high    LOG_SHIFT_VARIANT_COUNT    the shift variant
152/// ```
153///
154/// # Performance
155///
156/// Every entry is one copy or one accumulation.
157/// So the whole table costs `O(2^15)` field operations, and a single slice `O(2^6)`.
158/// A caller needing one slice rather than the whole table calls [`shift_operator_row`] itself.
159///
160/// # Panics
161///
162/// Panics unless the weights hold one entry per bit position of a word.
163#[instrument(skip_all, name = "shift_operator_table")]
164pub fn shift_operator_table<F, P: PackedField<Scalar = F>, A: Allocator>(
165	alloc: &A,
166	psi: &[F],
167) -> FieldVec<P, A>
168where
169	F: BinaryField,
170{
171	assert_eq!(psi.len(), Word::BITS, "the weights are indexed by bit position");
172	assert_eq!(
173		Word::BITS % P::WIDTH,
174		0,
175		"a row of Word::BITS weights must be packed-element aligned"
176	);
177
178	// One row of `Word::BITS` weights per `(variant, amount)`, variant most significant, packed
179	// `P::WIDTH` scalars at a time straight into the destination buffer. The scratch row is
180	// reused across every pair rather than collected into a second full-size buffer.
181	let row_packed_len = Word::BITS / P::WIDTH;
182	let packed_len = 1 << SHIFT_OPERATOR_LOG_LEN.saturating_sub(P::LOG_WIDTH);
183	let mut values = alloc.alloc::<P>(packed_len);
184	let mut row = [F::ZERO; Word::BITS];
185	for (variant, block) in iter::zip(
186		ShiftVariant::ALL,
187		values
188			.spare_capacity_mut()
189			.chunks_exact_mut(Word::BITS * row_packed_len),
190	) {
191		for (amount, packed_row) in block.chunks_exact_mut(row_packed_len).enumerate() {
192			shift_operator_row(variant, amount, &mut row, psi);
193			for (slot, chunk) in iter::zip(packed_row, row.chunks_exact(P::WIDTH)) {
194				slot.write(P::from_scalars(chunk.iter().copied()));
195			}
196		}
197	}
198	// Safety: the loop above wrote every one of the `packed_len` slots.
199	unsafe { values.set_len(packed_len) };
200
201	FieldBuffer::new(SHIFT_OPERATOR_LOG_LEN, values)
202}
203
204#[cfg(test)]
205mod tests {
206	use binius_compute::GlobalAllocator;
207	use binius_field::{Ghash128b, PackedGhash2x128b, Random, Rijndael8b};
208	use binius_math::{
209		BinarySubspace, inner_product::inner_product_buffers, multilinear::eq::eq_ind_partial_eval,
210		test_utils::random_scalars, univariate::EvaluationDomain,
211	};
212	use binius_verifier::protocols::shift::LOG_SHIFT_VARIANT_COUNT;
213	use proptest::prelude::*;
214	use rand::{SeedableRng, rngs::StdRng};
215
216	use super::{
217		super::{ShiftChallenge, ShiftChallengePoint, ShiftIndSumcheck},
218		*,
219	};
220
221	/// Phase 1's h multilinear and the claim phase 3 starts from must agree.
222	///
223	/// Phase 3 sums the shift indicators over the bit index, weighted by the Lagrange evaluations
224	/// and by the constant it carries; the multilinear holds those sums over the whole shift axis.
225	/// So with the carried constant set to one, evaluating the multilinear at `(r_j, r_s, r_v)`
226	/// must give phase 3's claim.
227	#[test]
228	fn h_op_consistency() {
229		type F = Ghash128b;
230		type P = PackedGhash2x128b;
231
232		let mut rng = StdRng::seed_from_u64(0);
233
234		let num_random_tests = 10;
235
236		for test_case in 0..num_random_tests {
237			let r_zhat_prime = F::random(&mut rng);
238
239			let r_j = random_scalars::<F>(&mut rng, Word::LOG_BITS);
240			let r_s = random_scalars::<F>(&mut rng, Word::LOG_BITS);
241			let r_v = random_scalars::<F>(&mut rng, LOG_SHIFT_VARIANT_COUNT);
242			let shift = ShiftChallenge::new(r_s.clone(), r_v.clone());
243
244			// Method 1: the claim phase 3 starts from, with the carried constant set to one.
245			let subspace = BinarySubspace::<Rijndael8b>::with_dim(Word::LOG_BITS).isomorphic();
246			let l_tilde = subspace.lagrange_evals_buffer(r_zhat_prime);
247			let claimed = ShiftIndSumcheck::<P, _>::new(
248				&GlobalAllocator,
249				l_tilde.as_ref(),
250				&ShiftChallengePoint::new(&r_j, &shift),
251				F::ONE,
252			)
253			.beta();
254
255			// Method 2: evaluate the built multilinear at the whole point.
256			let h = shift_operator_table::<F, P, _>(&GlobalAllocator, l_tilde.as_ref());
257			let evaluation_point = [r_j, r_s, r_v].concat();
258			let tensor = eq_ind_partial_eval::<P>(&evaluation_point);
259			let direct = inner_product_buffers(&h, &tensor);
260
261			assert_eq!(
262				claimed, direct,
263				"H-op evaluation mismatch (test_case={test_case}): claimed != direct",
264			);
265		}
266	}
267
268	/// Whether output bit `k` of `variant` at `amount` reads input bit `j`.
269	///
270	/// Read off the word operation itself.
271	/// Shifting a word with only bit `j` set leaves bits exactly where that bit is read.
272	fn reads_input_bit(variant: ShiftVariant, k: usize, j: usize, amount: usize) -> bool {
273		let shifted = variant.apply(Word(1u64 << j), amount);
274		(shifted.as_u64() >> k) & 1 == 1
275	}
276
277	/// The operator table computed straight from the indicator definition, one entry at a time.
278	fn reference_table<F: Field>(psi: &[F]) -> Vec<F> {
279		let mut table = vec![F::ZERO; 1 << SHIFT_OPERATOR_LOG_LEN];
280		for (variant_idx, variant) in ShiftVariant::ALL.into_iter().enumerate() {
281			for amount in 0..Word::BITS {
282				for j in 0..Word::BITS {
283					// Contract on the indicator's output-bit index, which is the summed one.
284					let entry = (0..Word::BITS)
285						.filter(|&k| reads_input_bit(variant, k, j, amount))
286						.map(|k| psi[k])
287						.sum();
288					table[(variant_idx * Word::BITS + amount) * Word::BITS + j] = entry;
289				}
290			}
291		}
292		table
293	}
294
295	proptest! {
296		// Invariant: every entry is the contraction the definition names.
297		//
298		// The eight variants and all 64 amounts are enumerated in full.
299		// Only the weights are sampled, since the operator is linear in them.
300		//
301		// This is what pins `sra`.
302		// Its vacated positions all read bit 63, so several weights land in one entry.
303		// That is the only slice which is not a plain move of the weights.
304		#[test]
305		fn shift_operator_table_matches_the_indicator_definition(seed: u64) {
306			type F = Ghash128b;
307
308			let mut rng = StdRng::seed_from_u64(seed);
309			let psi = random_scalars::<F>(&mut rng, Word::BITS);
310
311			let table = shift_operator_table::<F, F, _>(&GlobalAllocator, &psi);
312			let reference = reference_table(&psi);
313			prop_assert_eq!(table.as_ref(), reference.as_slice());
314		}
315
316		// Invariant: the operator is linear in the weights.
317		//
318		//     T[a * psi_1 + b * psi_2] == a * T[psi_1] + b * T[psi_2]
319		//
320		// The reduction folds the weights between its two contractions.
321		// Linearity is what lets it fold first and contract after.
322		#[test]
323		fn shift_operator_table_is_linear_in_the_weights(seed: u64) {
324			type F = Ghash128b;
325
326			let mut rng = StdRng::seed_from_u64(seed);
327			let psi_1 = random_scalars::<F>(&mut rng, Word::BITS);
328			let psi_2 = random_scalars::<F>(&mut rng, Word::BITS);
329			let (a, b) = (F::random(&mut rng), F::random(&mut rng));
330
331			// The combination pushed through the operator.
332			let combined = iter::zip(&psi_1, &psi_2)
333				.map(|(&x, &y)| a * x + b * y)
334				.collect::<Vec<F>>();
335			let lhs = shift_operator_table::<F, F, _>(&GlobalAllocator, &combined);
336
337			// The two tables combined afterwards.
338			let table_1 = shift_operator_table::<F, F, _>(&GlobalAllocator, &psi_1);
339			let table_2 = shift_operator_table::<F, F, _>(&GlobalAllocator, &psi_2);
340			let rhs = iter::zip(table_1.as_ref(), table_2.as_ref())
341				.map(|(&x, &y)| a * x + b * y)
342				.collect::<Vec<F>>();
343
344			prop_assert_eq!(lhs.as_ref(), rhs.as_slice());
345		}
346	}
347
348	// Invariant: the row builder writes exactly the slice the table holds for that pair.
349	//
350	// The reduction's outer phase reads one slice at a time rather than building the table, so the
351	// two paths have to agree entry for entry.
352	//
353	// The scratch row is carried across every pair and starts out non-zero, which is what pins the
354	// full-write contract: a builder that left cells alone would leak the previous pair's weights.
355	#[test]
356	fn shift_operator_row_matches_its_slice_of_the_table() {
357		type F = Ghash128b;
358
359		let mut rng = StdRng::seed_from_u64(0);
360		let psi = random_scalars::<F>(&mut rng, Word::BITS);
361
362		let table = shift_operator_table::<F, F, _>(&GlobalAllocator, &psi);
363		let mut row = vec![F::ONE; Word::BITS];
364		for (variant_idx, variant) in ShiftVariant::ALL.into_iter().enumerate() {
365			for amount in 0..Word::BITS {
366				shift_operator_row(variant, amount, &mut row, &psi);
367				let offset = (variant_idx * Word::BITS + amount) * Word::BITS;
368				assert_eq!(
369					row.as_slice(),
370					&table.as_ref()[offset..offset + Word::BITS],
371					"{variant:?} at amount {amount}"
372				);
373			}
374		}
375	}
376
377	// Invariant: the amount-zero slice hands back the weights untouched.
378	//
379	// A zero amount is the identity for every variant.
380	// That is what makes a single shift the special case of a sequence with a zero outer amount.
381	//
382	// The half-word forms read the amount modulo 32, so they are the identity at zero as well.
383	#[test]
384	fn the_zero_amount_slice_returns_the_weights_unchanged() {
385		type F = Ghash128b;
386
387		let mut rng = StdRng::seed_from_u64(0);
388		let psi = random_scalars::<F>(&mut rng, Word::BITS);
389
390		let table = shift_operator_table::<F, F, _>(&GlobalAllocator, &psi);
391		for (variant_idx, variant) in ShiftVariant::ALL.into_iter().enumerate() {
392			// Amount zero sits at the front of the variant's block of rows.
393			let row = &table.as_ref()[variant_idx * Word::BITS * Word::BITS..][..Word::BITS];
394			assert_eq!(row, psi.as_slice(), "{variant:?} at amount zero is not the identity");
395		}
396	}
397}