Skip to main content

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}