binius_prover/protocols/shift/claims.rs
1// Copyright 2026 The Binius Developers
2
3//! The shift reduction's operand evaluation claims, all at one constraint point.
4//!
5//! The reduction closes four constraint families in one proof: ZERO, AND, IMUL and BMUL.
6//!
7//! Their operands form one flat run of columns, in [`OPERATION_ARITIES`] order.
8//! Each column arrives with its evaluation, claimed at the one constraint point `r_x` with its
9//! matrix padded with empty rows. Those claims then travel together through both phases of the
10//! reduction, batched on one operand axis.
11
12use std::iter;
13
14use binius_core::constraint_system::ConstraintSystem;
15use binius_field::Field;
16use binius_ip_prover::channel::IPProverChannel;
17use binius_math::{
18 inner_product::inner_product,
19 multilinear::eq::{eq_ind_partial_eval, eq_ind_partial_eval_scalars},
20};
21use binius_utils::checked_arithmetics::log2_ceil_usize;
22use binius_verifier::{
23 protocols::{rerand::RerandOutput, zero},
24 reduction::{BINMUL_ARITY, INTMUL_ARITY, OPERATION_ARITIES, ZERO_ARITY, padding_scales},
25};
26
27/// The operand evaluation claims of every operation, as the shift reduction receives them.
28#[derive(Debug, Clone)]
29pub struct OperandClaims<F: Field> {
30 /// The constraint point, as long as the widest operation's constraint count.
31 ///
32 /// Every operation is claimed at the whole point, its matrix padded with empty rows; see
33 /// [`padding_scales`].
34 pub r_x: Vec<F>,
35 /// The univariate challenge folding the bit axis, shared by every operation.
36 pub r_zhat_prime: F,
37 /// One evaluation per operand column, the four operations' runs in [`OPERATION_ARITIES`]
38 /// order.
39 pub evals: Vec<F>,
40}
41
42impl<F: Field> OperandClaims<F> {
43 /// Assembles the claims from the output of the BitAnd sumcheck.
44 ///
45 /// The sumcheck's point is `r_rho || r_x_star`, instance index low, constraint index high.
46 /// `r_x_star` spans the widest AND, IMUL and BMUL set. The constraint point `r_x` extends it
47 /// through [`zero::reduction_point`] when the ZERO set is wider still, drawing the extra
48 /// challenges from `sample`.
49 ///
50 /// The sumcheck claims each operation at its own prefix of `r_x`. Its padding factor lifts the
51 /// claim to its padded matrix at the whole point.
52 ///
53 /// An operation the constraint system does not use gets zero evaluations.
54 ///
55 /// # Arguments
56 ///
57 /// - `cs`: the constraint system, whose row counts pick each operation's padding factor.
58 /// - `log_instances`: the instance variables at the low end of the point.
59 /// - `z_challenge`: the univariate challenge, shared by every operation.
60 /// - `rerand`: the sumcheck's output, with the IntMul operand evaluations before the BinMul
61 /// ones.
62 /// - `sample`: draws the constraint point's extra challenges from the transcript.
63 pub fn from_rerand(
64 cs: &ConstraintSystem,
65 log_instances: usize,
66 z_challenge: F,
67 rerand: &RerandOutput<F>,
68 sample: impl FnMut() -> F,
69 ) -> Self {
70 let r_x_star = &rerand.eval_point[log_instances..];
71 let log_n_zero = cs.log_zero_constraints().unwrap_or(0);
72 let r_x = zero::reduction_point(r_x_star, r_x_star.len().max(log_n_zero), sample);
73
74 let [_, bitand_scale, intmul_scale, binmul_scale] = padding_scales(cs, &r_x);
75 let mut operand_evals = rerand.operand_evals.iter().copied();
76 // An absent operation's run is zeros; a present one takes its evaluations off the front.
77 let mut run = |present: bool, arity: usize, scale: F| {
78 if present {
79 operand_evals
80 .by_ref()
81 .take(arity)
82 .map(|eval| eval * scale)
83 .collect()
84 } else {
85 vec![F::ZERO; arity]
86 }
87 };
88 let intmul = run(cs.n_imul_constraints() > 0, INTMUL_ARITY, intmul_scale);
89 let binmul = run(cs.n_bmul_constraints() > 0, BINMUL_ARITY, binmul_scale);
90
91 // The BitAnd check has no skip branch: an empty AND set reduces over one zero row.
92 let bitand = rerand.bitand_evals.map(|eval| eval * bitand_scale);
93 let evals = iter::repeat_n(F::ZERO, ZERO_ARITY)
94 .chain(bitand)
95 .chain(intmul)
96 .chain(binmul)
97 .collect::<Vec<_>>();
98 debug_assert_eq!(evals.len(), OPERATION_ARITIES.iter().sum::<usize>());
99
100 Self {
101 r_x,
102 r_zhat_prime: z_challenge,
103 evals,
104 }
105 }
106
107 /// Draws the batching challenges and folds their weights into the claims.
108 ///
109 /// The claim of column `m` is weighted by the equality indicator of the operand axis at `m`.
110 /// The axis is padded to a cube of `log2_ceil(evals.len())` challenges; the columns past the
111 /// last claim name nothing. The verifier draws the same challenges at the same place, so the
112 /// count is protocol, not detail.
113 ///
114 /// # Arguments
115 ///
116 /// - `channel`: the transcript the batching challenges are drawn from.
117 pub fn prepare(self, channel: &mut impl IPProverChannel<F>) -> PreparedOperandClaims<F> {
118 let operand_batch_challenges = channel.sample_many(log2_ceil_usize(self.evals.len()));
119 let operand_weights = eq_ind_partial_eval_scalars(&operand_batch_challenges);
120 let batched_eval = inner_product(
121 self.evals.iter().copied(),
122 operand_weights[..self.evals.len()].iter().copied(),
123 );
124
125 PreparedOperandClaims {
126 batched_eval,
127 r_zhat_prime: self.r_zhat_prime,
128 r_x_tensor: eq_ind_partial_eval::<F>(&self.r_x).into_inner(),
129 operand_weights,
130 }
131 }
132}
133
134/// The claims with their batching weights folded in, as both proving phases read them.
135///
136/// The constraint table is shared by every key, so it is built once here.
137#[derive(Debug, Clone)]
138pub struct PreparedOperandClaims<F: Field> {
139 /// The operand claims collapsed into the single value the reduction proves:
140 ///
141 /// ```text
142 /// batched_eval = sum_m evals[m] * operand_weights[m]
143 /// ```
144 ///
145 /// This is the claim phase 1 hands its sumcheck, and it is the same value the verifier
146 /// computes from the operand evaluation claims before running its own.
147 pub batched_eval: F,
148 /// The univariate challenge folding the bit axis, shared by every operation.
149 pub r_zhat_prime: F,
150 /// The constraint table: the equality indicator of `r_x`, one weight per row of the padded
151 /// operation matrices, shared by every operation.
152 pub r_x_tensor: Vec<F>,
153 /// The weight of each operand column, padded to a power of two.
154 ///
155 /// A key reads it at the column its constraint index names.
156 pub operand_weights: Vec<F>,
157}
158
159#[cfg(test)]
160mod tests {
161 use binius_transcript::ProverTranscript;
162 use binius_verifier::config::{B128, StdChallenger};
163
164 use super::*;
165
166 #[test]
167 fn prepare_draws_one_challenge_per_operand_axis_variable() {
168 // Invariant: the operand axis takes `log2_ceil(n)` challenges, drawn first, and its
169 // weights batch the claims.
170 //
171 // The verifier draws in that order; any other weights the claims by a different tensor.
172 //
173 // Two transcripts from the same seed hand out the same sequence, so drawing the axis by
174 // hand from one pins what `prepare` must have drawn from the other.
175 let n = OPERATION_ARITIES.iter().sum::<usize>();
176 let evals = (1..=n as u128).map(B128::new).collect::<Vec<_>>();
177
178 let mut expected = ProverTranscript::<StdChallenger>::default();
179 let operand_weights =
180 eq_ind_partial_eval_scalars::<B128>(&expected.sample_many(log2_ceil_usize(n)));
181
182 let mut channel = ProverTranscript::<StdChallenger>::default();
183 let prepared = OperandClaims {
184 r_x: Vec::new(),
185 r_zhat_prime: B128::ZERO,
186 evals: evals.clone(),
187 }
188 .prepare(&mut channel);
189
190 // The two stay in lockstep only if `prepare` drew exactly those and no more.
191 assert_eq!(
192 IPProverChannel::<B128>::sample(&mut channel),
193 IPProverChannel::<B128>::sample(&mut expected)
194 );
195
196 assert_eq!(prepared.operand_weights, operand_weights);
197 assert_eq!(prepared.operand_weights.len(), 16);
198 assert_eq!(
199 prepared.batched_eval,
200 inner_product(evals, operand_weights[..n].iter().copied())
201 );
202 // The constraint point is empty, so the constraint table is the single weight one.
203 assert_eq!(prepared.r_x_tensor, [B128::ONE]);
204 }
205}