binius_prover/protocols/shift/
monster.rs1use 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
14const HALF_WORD_BITS: usize = 32;
16
17pub(super) struct OuterSlotWeights<F: Field> {
23 variant: FieldBuffer<F>,
25 amount: FieldBuffer<F>,
27}
28
29impl<F: Field> OuterSlotWeights<F> {
30 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 #[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
45pub 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 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 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 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#[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 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 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 #[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 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 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 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 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 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 #[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 #[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 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 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 #[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 #[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 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}