Skip to main content

binius_prover/protocols/shift/
outer.rs

1// Copyright 2026 The Binius Developers
2
3//! The rounds binding the outer slot of a shift sequence.
4//!
5//! A shifted value index names a word together with two shifts applied in sequence, and the shift
6//! reduction peels them from the output end inward. These are the rounds that peel the outer one.
7
8use std::iter;
9
10use binius_compute::Allocator;
11use binius_core::{ShiftVariant, word::Word};
12use binius_field::{BinaryField, Field, PackedField};
13use binius_ip::sumcheck::RoundCoeffs;
14use binius_ip_prover::sumcheck::round_evals::RoundEvals;
15use binius_math::{FieldVec, multilinear::fold::fold_highest_var_inplace};
16use binius_verifier::protocols::shift::{LOG_SHIFT_COUNT, SHIFT_COUNT};
17
18use super::{
19	monster::{shift_operator_row, shift_operator_table},
20	phase_1::SparseShiftRows,
21};
22
23/// The `(variant, amount)` pair a shift index names.
24///
25/// This inverts [`Shift::index`](binius_core::constraint_system::Shift::index) over the reduction's
26/// index space rather than over well-formed shifts: the amount axis spans `Word::BITS` for every
27/// variant, so a half-word (`*32`) variant's index may carry an amount no `Shift` of that variant
28/// could hold. Such an index is a hypercube vertex the sumcheck ranges over all the same.
29///
30/// # Panics
31///
32/// Panics unless the index is below [`SHIFT_COUNT`].
33pub fn decode_shift(index: usize) -> (ShiftVariant, usize) {
34	assert!(index < SHIFT_COUNT, "a shift index names one slot's spelling");
35	let variant = ShiftVariant::from_u8((index >> Word::LOG_BITS) as u8)
36		.expect("an index below SHIFT_COUNT has a variant field below the variant count");
37	(variant, index % Word::BITS)
38}
39
40/// The sumcheck rounds binding the outer shift of a sequence, against the folded oblong table.
41///
42/// The reduction's `h` factor spans 24 variables — the bit position and both shift slots — so its
43/// value table would hold `2^24` entries. It is never formed. Instead, writing `T` for the shift
44/// operator ([`shift_operator_table`]) and `d` for the oblong weights,
45///
46/// ```text
47///     eta := T[d]                                    2^15 entries
48///     h(J, s_2, o_2, s_1, o_1) = T[eta(., s_2, o_2)](J, s_1, o_1)
49/// ```
50///
51/// so `h` is reached by applying `T` to the oblong weights, taking the outer slice, and applying
52/// `T` again. This stage holds `eta` and folds it as the rounds bind the outer slot; each round's
53/// `h` rows are derived from the folded table, one slice at a time. Folding `eta` first and
54/// applying `T` after gives the same answer as folding `h`, because `T` is linear in its weights.
55///
56/// # Why the outer slot binds first
57///
58/// Two independent reasons:
59///
60/// - **Correctness.** The two indicator matrices do not commute — `sra` is the obstruction, since
61///   it is the one shift whose vacated positions all read a single input bit. Nesting `T` inside
62///   `T` composes them in the order the slots are bound, so only binding the outer slot first
63///   computes `h` rather than its transpose-order counterpart.
64/// - **Cost.** The push-through is a *shift* only while the inner pair is still a cube index. Under
65///   the opposite order the outer indicator would arrive folded to a dense `2^6 x 2^6` matrix, and
66///   each live shift quadruple would cost a matrix-vector product in place of a shift.
67pub struct OuterShiftStage<F: Field, A: Allocator> {
68	/// `eta`, the oblong weights pushed through one shift, folded over the outer variables bound
69	/// so far.
70	///
71	/// The axes run, from the low index positions up: the intermediate bit index, then the outer
72	/// shift amount, then the outer shift variant. So binding the highest variable takes the outer
73	/// variant before the outer amount, which is the order the reduction's rounds run in.
74	eta: FieldVec<F, A>,
75}
76
77impl<F: BinaryField, A: Allocator> OuterShiftStage<F, A> {
78	/// Pushes the oblong weights through every shift, ready for the first round.
79	///
80	/// # Panics
81	///
82	/// Panics unless the weights hold one entry per bit position of a word.
83	pub fn new(alloc: &A, oblong_weights: &[F]) -> Self {
84		Self {
85			eta: shift_operator_table::<F, F, A>(alloc, oblong_weights),
86		}
87	}
88
89	/// The number of outer-index variables the stage has yet to bind.
90	pub const fn n_vars_remaining(&self) -> usize {
91		self.eta.log_len() - Word::LOG_BITS
92	}
93
94	/// The folded table at one outer index: `eta(., outer)` over the intermediate bit index.
95	fn stride(&self, outer: usize) -> &[F] {
96		&self.eta.as_ref()[outer * Word::BITS..][..Word::BITS]
97	}
98
99	/// The weights the inner rounds run against: `eta` at the bound outer point.
100	///
101	/// The terminal fold of `eta` *is* the partial evaluation
102	/// `sum_i d(i) * shift-ind~(i, K, r_s2, r_o2)`, the multilinear extension commuting with the
103	/// finite sum over `i`. So the inner rounds need no division and no second pass.
104	///
105	/// ## Preconditions
106	///
107	/// * `self.n_vars_remaining() == 0`
108	pub fn psi(&self) -> &[F] {
109		assert_eq!(self.n_vars_remaining(), 0, "precondition: every outer variable is bound");
110		self.eta.as_ref()
111	}
112
113	/// Binds the highest outer variable to a challenge.
114	///
115	/// ## Preconditions
116	///
117	/// * `self.n_vars_remaining() >= 1`
118	pub fn fold(&mut self, challenge: F) {
119		assert!(self.n_vars_remaining() > 0, "precondition: an outer variable remains to bind");
120		fold_highest_var_inplace(&mut self.eta, challenge);
121	}
122
123	/// Computes one round message: the degree-2 round polynomial binding the next outer variable.
124	///
125	/// The round polynomial is sampled at 1 and at infinity, as [`RoundEvals`] documents; the claim
126	/// supplies its value at 0. Both are linear in `g`, so each stored row contributes on its own
127	/// and rows facing each other across the split never have to be paired:
128	///
129	/// ```text
130	/// R(1)   = sum_v G_1(v) H_1(v)             row (i, c) adds <c, h[i]>, upper half only
131	/// R(inf) = sum_v (G_0 + G_1)(H_0 + H_1)    row (i, c) adds <c, h[i] + h[i ^ half]>, either half
132	/// ```
133	///
134	/// Unlike the rounds that follow, the `h` rows are not read from a table but derived: a row's
135	/// is one slice of the shift operator applied to one stride of the folded `eta`, which costs
136	/// `O(2^6)`. A round therefore costs `O(2^6 * n_shift)` in the number of live shift quadruples,
137	/// with no charge proportional to the space they are drawn from.
138	///
139	/// ## Preconditions
140	///
141	/// * `g`'s row index is a shift quadruple, the outer slot above the inner one
142	/// * `g` and this stage have the same number of outer variables left to bind
143	pub fn round_coeffs<P: PackedField<Scalar = F>>(
144		&self,
145		g: &SparseShiftRows<P>,
146		claim: F,
147	) -> RoundCoeffs<F> {
148		assert_eq!(
149			g.log_rows(),
150			self.n_vars_remaining() + LOG_SHIFT_COUNT,
151			"precondition: the rows and the folded table agree on the outer index"
152		);
153
154		// The bit this round binds is an outer one, since the outer slot sits above the inner one
155		// in a quadruple. Dropping the inner slot off it leaves the bit that indexes eta's strides.
156		let half = g.half();
157		let facing_half = half >> LOG_SHIFT_COUNT;
158
159		// One scratch row per side, rewritten for each stored row: `shift_operator_row` writes
160		// every cell, so nothing carries over between rows and neither needs an allocation.
161		let mut own = [F::ZERO; Word::BITS];
162		let mut facing = [F::ZERO; Word::BITS];
163
164		let (mut y_1, mut y_inf) = (F::ZERO, F::ZERO);
165		for (index, row) in g.rows() {
166			// A row and the row facing it share an inner slot and differ in the outer bit being
167			// bound, so one slice of the operator serves both.
168			let (variant, amount) = decode_shift(index % SHIFT_COUNT);
169			let outer = index >> LOG_SHIFT_COUNT;
170			shift_operator_row(variant, amount, &mut own, self.stride(outer));
171			shift_operator_row(variant, amount, &mut facing, self.stride(outer ^ facing_half));
172
173			// The infinity evaluation reads H(0) + H(1), the same sum from either half, so the
174			// row's own half only decides the evaluation at 1.
175			let in_upper_half = index & half != 0;
176			for (value, (&own_j, &facing_j)) in
177				iter::zip(P::iter_slice(row), iter::zip(&own, &facing))
178			{
179				if in_upper_half {
180					y_1 += value * own_j;
181				}
182				y_inf += value * (own_j + facing_j);
183			}
184		}
185
186		RoundEvals([y_1, y_inf]).interpolate(claim)
187	}
188}
189
190#[cfg(test)]
191mod tests {
192	use binius_compute::GlobalAllocator;
193	use binius_core::constraint_system::Shift;
194	use binius_field::{Ghash128b, Random};
195	use binius_math::test_utils::random_scalars;
196	use rand::{SeedableRng, rngs::StdRng};
197
198	use super::*;
199
200	type F = Ghash128b;
201
202	/// Whether output bit `out` of `variant` at `amount` reads input bit `in_bit`.
203	///
204	/// Read off the word operation itself: shifting a word with only bit `in_bit` set leaves bits
205	/// exactly where that bit is read. This is the shift indicator, straight from its definition
206	/// and independent of every table the reduction builds.
207	fn reads_input_bit(variant: ShiftVariant, out: usize, in_bit: usize, amount: usize) -> bool {
208		let shifted = variant.apply(binius_core::word::Word(1u64 << in_bit), amount);
209		(shifted.as_u64() >> out) & 1 == 1
210	}
211
212	/// The `h` rows of one inner shift, over every outer index, straight from the definition:
213	///
214	/// ```text
215	/// h(j) = sum_{i, k} d(i) * shift-ind(i, k, outer) * shift-ind(k, j, inner)
216	/// ```
217	///
218	/// This is the double contraction the stage computes by nesting the shift operator. It is
219	/// evaluated only where `g` is supported, so no `2^24` table is ever formed — and no
220	/// multiplications are needed, since each indicator selects rather than scales.
221	fn reference_rows(d: &[F], inner: Shift) -> Vec<Vec<F>> {
222		let (inner_variant, inner_amount) = (inner.variant, inner.amount as usize);
223		(0..SHIFT_COUNT)
224			.map(|outer_index| {
225				let (outer_variant, outer_amount) = decode_shift(outer_index);
226				// eta at this outer index: the oblong weights carried to the intermediate word.
227				let eta = (0..Word::BITS)
228					.map(|k| {
229						(0..Word::BITS)
230							.filter(|&i| reads_input_bit(outer_variant, i, k, outer_amount))
231							.map(|i| d[i])
232							.sum::<F>()
233					})
234					.collect::<Vec<F>>();
235				// And on down to the witness bit.
236				(0..Word::BITS)
237					.map(|j| {
238						(0..Word::BITS)
239							.filter(|&k| reads_input_bit(inner_variant, k, j, inner_amount))
240							.map(|k| eta[k])
241							.sum::<F>()
242					})
243					.collect()
244			})
245			.collect()
246	}
247
248	/// A dense reference for the rounds, folded directly rather than derived from a folded `eta`.
249	///
250	/// One `(g, h)` pair per inner shift the fixture uses, each spanning the whole outer index. The
251	/// outer rounds never mix inner indices and `g` is zero at the inner shifts absent here, so
252	/// this is exact rather than a restriction.
253	struct Reference {
254		/// Per inner shift, the `g` rows over the remaining outer index.
255		g: Vec<Vec<Vec<F>>>,
256		/// Per inner shift, the `h` rows over the remaining outer index.
257		h: Vec<Vec<Vec<F>>>,
258	}
259
260	impl Reference {
261		/// The sum the rounds start from.
262		fn sum(&self) -> F {
263			iter::zip(&self.g, &self.h)
264				.flat_map(|(g, h)| iter::zip(g, h))
265				.flat_map(|(g_row, h_row)| iter::zip(g_row, h_row))
266				.map(|(&g, &h)| g * h)
267				.sum()
268		}
269
270		/// This round's message, by brute force over every entry.
271		fn round_coeffs(&self, claim: F) -> RoundCoeffs<F> {
272			let half = self.g[0].len() / 2;
273			let (mut y_1, mut y_inf) = (F::ZERO, F::ZERO);
274			for (g, h) in iter::zip(&self.g, &self.h) {
275				for lower in 0..half {
276					let upper = lower + half;
277					for j in 0..Word::BITS {
278						y_1 += g[upper][j] * h[upper][j];
279						y_inf += (g[lower][j] + g[upper][j]) * (h[lower][j] + h[upper][j]);
280					}
281				}
282			}
283			RoundEvals([y_1, y_inf]).interpolate(claim)
284		}
285
286		/// Binds the highest outer variable of both tables.
287		fn fold(&mut self, challenge: F) {
288			let fold = |rows: &mut Vec<Vec<F>>| {
289				let half = rows.len() / 2;
290				for lower in 0..half {
291					for j in 0..Word::BITS {
292						let (low, high) = (rows[lower][j], rows[lower + half][j]);
293						rows[lower][j] = low + challenge * (low + high);
294					}
295				}
296				rows.truncate(half);
297			};
298			self.g.iter_mut().for_each(&fold);
299			self.h.iter_mut().for_each(&fold);
300		}
301	}
302
303	/// The inner shifts of the fixture, one `g` row per outer index the fixture stores them at.
304	///
305	/// Every variant appears, and every case whose operator slice is not a plain move of the
306	/// weights: `sra` and `sra32` pile several weights onto one position, in either slot.
307	fn fixture() -> (Vec<F>, Vec<(Shift, Vec<Shift>)>) {
308		let mut rng = StdRng::seed_from_u64(0);
309		let d = random_scalars::<F>(&mut rng, Word::BITS);
310		let quadruples = vec![
311			// A sign extension: shift a field up to the top and arithmetically back down.
312			(Shift::sll(40), vec![Shift::sar(40)]),
313			// The sign bit in the inner slot instead, under two different outer shifts.
314			(Shift::sar(7), vec![Shift::srl(3), Shift::rotr(19)]),
315			// A rotate under a rotate, which wraps from both ends.
316			(Shift::rotr(1), vec![Shift::rotr(63), Shift::sll(9)]),
317			// The half-word family, including its own sign-extension case.
318			(Shift::sll32(11), vec![Shift::sra32(11), Shift::rotr32(5)]),
319			(Shift::srl32(3), vec![Shift::sll32(30)]),
320			// An unshifted inner slot, which the reduction reaches as the identity spelling.
321			(Shift::IDENTITY, vec![Shift::IDENTITY, Shift::srl(17)]),
322		];
323		(d, quadruples)
324	}
325
326	/// The stage's round messages are the ones a prover folding `h` itself would send.
327	///
328	/// This is the property the whole construction rests on: deriving each row from a folded `eta`
329	/// gives the same round polynomial as folding the fully formed `h`, round after round, because
330	/// the shift operator is linear in its weights.
331	#[test]
332	fn round_messages_match_a_directly_folded_reference() {
333		let mut rng = StdRng::seed_from_u64(1);
334		let (d, quadruples) = fixture();
335
336		// The sparse rows, keyed on the quadruple with the outer slot above the inner one.
337		let mut indices = Vec::new();
338		let mut values = Vec::new();
339		let mut reference_g = Vec::new();
340		for (inner, outers) in &quadruples {
341			let mut rows = vec![vec![F::ZERO; Word::BITS]; SHIFT_COUNT];
342			for outer in outers {
343				let row = random_scalars::<F>(&mut rng, Word::BITS);
344				indices.push((outer.index() << LOG_SHIFT_COUNT | inner.index()) as u32);
345				values.extend_from_slice(&row);
346				// Rows at a repeated index add up, which the reference has to mirror.
347				for (slot, value) in iter::zip(&mut rows[outer.index()], row) {
348					*slot += value;
349				}
350			}
351			reference_g.push(rows);
352		}
353		let log_rows = 2 * LOG_SHIFT_COUNT;
354		let mut g = SparseShiftRows::<F>::new(indices, values, log_rows);
355
356		let mut reference = Reference {
357			g: reference_g,
358			h: quadruples
359				.iter()
360				.map(|(inner, _)| reference_rows(&d, *inner))
361				.collect(),
362		};
363
364		let mut stage = OuterShiftStage::new(&GlobalAllocator, &d);
365		assert_eq!(stage.n_vars_remaining(), LOG_SHIFT_COUNT);
366
367		// A non-degenerate fixture, or matching round messages would prove nothing.
368		let mut claim = reference.sum();
369		assert_ne!(claim, F::ZERO);
370
371		for _ in 0..LOG_SHIFT_COUNT {
372			let coeffs = stage.round_coeffs(&g, claim);
373			assert_eq!(coeffs, reference.round_coeffs(claim));
374
375			let challenge = F::random(&mut rng);
376			claim = coeffs.evaluate(&challenge);
377			stage.fold(challenge);
378			g.fold(challenge);
379			reference.fold(challenge);
380		}
381
382		// The stage hands the inner rounds `eta` at the bound outer point, which is what the
383		// reference's own `h` was folded down to for the inner shift that leaves it untouched.
384		assert_eq!(stage.n_vars_remaining(), 0);
385		let identity_slot = quadruples
386			.iter()
387			.position(|(inner, _)| *inner == Shift::IDENTITY)
388			.expect("the fixture carries an identity inner slot");
389		assert_eq!(stage.psi(), reference.h[identity_slot][0].as_slice());
390	}
391
392	/// The stage never reads a table indexed by anything but the intermediate bit and the outer
393	/// slot, so its cost is fixed however many quadruples are live.
394	#[test]
395	fn the_folded_table_stays_the_size_of_one_shift_slot() {
396		let (d, _) = fixture();
397		let mut stage = OuterShiftStage::<F, _>::new(&GlobalAllocator, &d);
398
399		let mut expected = Word::LOG_BITS + LOG_SHIFT_COUNT;
400		assert_eq!(stage.eta.log_len(), expected);
401		for _ in 0..LOG_SHIFT_COUNT {
402			stage.fold(F::ONE);
403			expected -= 1;
404			assert_eq!(stage.eta.log_len(), expected);
405		}
406		assert_eq!(stage.psi().len(), Word::BITS);
407	}
408}