1use 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
21const HALF_WORD_LOG_BITS: usize = Word::LOG_BITS - 1;
23
24#[derive(Debug, Clone)]
26pub struct ShiftChallenge<F> {
27 pub(crate) amount: Vec<F>,
29 pub(crate) variant: Vec<F>,
31}
32
33impl<F> ShiftChallenge<F> {
34 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#[derive(Debug, Clone, Copy)]
44pub struct ShiftChallengePoint<'a, F> {
45 bit: &'a [F],
47 shift: &'a ShiftChallenge<F>,
49}
50
51impl<'a, F: BinaryField> ShiftChallengePoint<'a, F> {
52 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 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 let sign_position: F = bit.iter().copied().product();
71
72 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 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
109pub struct ShiftIndSumcheck<P: PackedField, A: Allocator> {
134 scaled_weights: FieldVec<P, A>,
136 shift_ind: FieldVec<P, A>,
138 weights: Vec<P::Scalar>,
140 beta: P::Scalar,
142}
143
144#[derive(Debug, Clone)]
151pub struct ShiftIndOutput<F> {
152 pub weights_eval: F,
154 pub ind_eval: F,
157 pub eval: F,
159 pub point: Vec<F>,
161}
162
163impl<F: BinaryField, P: PackedField<Scalar = F>, A: Allocator> ShiftIndSumcheck<P, A> {
164 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 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 pub const fn beta(&self) -> F {
201 self.beta
202 }
203
204 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 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
237fn 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 for k in 0..n {
252 let j_k = bit[k].clone();
253 let s_k = amount[k].clone();
254
255 let both = j_k.clone() * &s_k;
257 let j_one_s = j_k.clone() - &both; let one_j_s = s_k.clone() - &both; let xor = j_k + s_k;
260 let eq = E::one() + &xor;
261
262 for i in 0..(1 << k) {
264 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 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
279fn 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 for k in 0..n {
288 let s_k = amount[k].clone();
289
290 for i in 0..(1 << k) {
292 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
302fn 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 for k in 0..n {
315 let j_k = bit[k].clone();
316 let s_k = amount[k].clone();
317
318 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 for i in 0..(1 << k) {
326 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 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 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 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 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 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 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 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 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 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 #[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 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 #[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}