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}