Skip to main content

binius_prover/protocols/shift/
phase_2.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{cmp::max, iter};
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, PackedField, WideMul};
9use binius_ip::sumcheck::{RoundCoeffs, SumcheckOutput};
10use binius_ip_prover::{
11	channel::IPProverChannel,
12	sumcheck::{
13		ProveSingleOutput, bivariate_product_prover, prove_single, round_evals::RoundEvals,
14	},
15};
16use binius_math::{
17	FieldVec,
18	multilinear::eq::{eq_ind_partial_eval, eq_ind_zero},
19};
20use binius_utils::{
21	checked_arithmetics::log2_ceil_usize,
22	rayon::{
23		prelude::*,
24		task_size::{IndexedParallelIteratorExt, WorkPerItem},
25	},
26};
27use binius_verifier::protocols::shift::evaluate_words_mle;
28use tracing::instrument;
29
30use super::{
31	SegmentWords, claims::PreparedOperandClaims, key_collection::KeyCollection,
32	phase_1::Phase1Output,
33};
34use crate::fold_word::BitAxisFolder;
35
36/// Proves the second phase of the shift protocol reduction.
37///
38/// Folds the value-vector words by the bit-position challenge.
39/// Builds the constraint-matrix multilinear's two segments.
40/// Then runs a sumcheck between them, with a sparse first round over the segment selector.
41///
42/// # Arguments
43///
44/// - `key_collection`: the prover's key collection for the constraint system.
45/// - `words`: the value-vector words.
46/// - `prepared`: the prepared claim of each operation, indexed by the operation a key names.
47/// - `phase_1_output`: the challenges and evaluation the first phase produced.
48/// - `shift_ind_eval`: the scalar weighting every shift key.
49/// - `epsilon`: the claim this phase's rounds prove.
50/// - `channel`: the prover channel the interactive rounds run over.
51/// - `alloc`: the allocator the intermediate buffers are drawn from.
52///
53/// `shift_ind_eval` is the product of the two indicator evaluations.
54/// Those are what the earlier bit-index phases reduced to.
55///
56/// # Returns
57///
58/// The combined challenges with the witness evaluation, and the wiring multilinear's evaluation.
59#[allow(clippy::too_many_arguments)]
60#[instrument(skip_all, name = "prove_phase_2")]
61pub fn prove_phase_2<F, P, Channel, A>(
62	key_collection: &KeyCollection,
63	words: SegmentWords<'_>,
64	prepared: &PreparedOperandClaims<F>,
65	phase_1_output: Phase1Output<F>,
66	shift_ind_eval: F,
67	epsilon: F,
68	channel: &mut Channel,
69	alloc: &A,
70) -> ShiftOutput<F>
71where
72	F: BinaryField,
73	P: PackedField<Scalar = F>,
74	Channel: IPProverChannel<F>,
75	A: Allocator,
76{
77	let Phase1Output {
78		r_j,
79		inner,
80		outer,
81		psi: _,
82		gamma: _,
83		g_eval: _,
84	} = phase_1_output;
85
86	let r_j_tensor = eq_ind_partial_eval::<F>(&r_j);
87
88	// Fold each segment separately.
89	// The combined witness is never materialized.
90	// Each fold is zero-padded to enough variables to cover its own segment's length.
91	// Both columns fold against the same round tensor, so the tables are built once.
92	let folder = BitAxisFolder::new(r_j_tensor.as_ref());
93	let public_folded = folder.fold::<P, _>(alloc, words.public);
94	let hidden_folded = folder.fold::<P, _>(alloc, words.hidden);
95
96	let (public_monster, hidden_monster) =
97		key_collection.build_monster_segments(alloc, prepared, shift_ind_eval, &inner, &outer);
98
99	// Both halves of the sumcheck share one word-index space, spanning the wider segment.
100	// The hidden segment is normally the wider one.
101	// A system with more public words than private values inverts that.
102	// So the hidden half is zero-extended to match.
103	let log_segment_words = max(public_folded.log_len(), hidden_folded.log_len());
104	let hidden_folded = hidden_folded.zero_extend_in(alloc, log_segment_words);
105	let hidden_monster = hidden_monster.zero_extend_in(alloc, log_segment_words);
106
107	run_sumcheck(
108		&public_folded,
109		hidden_folded,
110		&public_monster,
111		hidden_monster,
112		shift_ind_eval,
113		words.public,
114		r_j,
115		epsilon,
116		channel,
117		alloc,
118	)
119}
120
121/// A witness or constraint-matrix buffer, split into the public and hidden segments the
122/// phase-2 sumcheck's selector variable chooses between.
123///
124/// The hidden segment is normally the wider of the two, spanning the whole word-index space,
125/// with the public segment sitting at its base.
126struct SegmentPair<'a, P: PackedField, A: Allocator> {
127	/// The public segment, at the base of the shared word-index space.
128	public: &'a FieldVec<P, A>,
129	/// The hidden segment, normally spanning the whole word-index space.
130	hidden: FieldVec<P, A>,
131}
132
133impl<'a, P: PackedField, A: Allocator> SegmentPair<'a, P, A> {
134	/// Pairs a public and hidden segment sharing one word-index space.
135	const fn new(public: &'a FieldVec<P, A>, hidden: FieldVec<P, A>) -> Self {
136		Self { public, hidden }
137	}
138
139	/// Folds the two segments at the selector challenge.
140	///
141	/// Consumes and overwrites the hidden buffer for memory efficiency.
142	/// The result is `(1 - alpha) * public_padded + alpha * hidden`, exactly what folding the
143	/// materialized combined buffer's highest variable would produce.
144	fn fold<F>(self, alpha: F) -> FieldVec<P, A>
145	where
146		F: Field,
147		P: PackedField<Scalar = F>,
148	{
149		let Self { public, mut hidden } = self;
150
151		// Scale the dominant hidden segment in place, in parallel.
152		let alpha_broadcast = P::broadcast(alpha);
153		hidden
154			.as_mut()
155			.par_iter_mut()
156			.with_min_task(WorkPerItem::FieldMuls)
157			.for_each(|hidden_i| *hidden_i *= alpha_broadcast);
158
159		// Add the small public prefix sequentially.
160		// Its trailing partial packed element carries zero high lanes, so whole-element
161		// updates are correct.
162		let one_minus_alpha = P::broadcast(F::ONE - alpha);
163		let n_public_packed = public.as_ref().len();
164		for (value, &public_i) in
165			iter::zip(&mut hidden.as_mut()[..n_public_packed], public.as_ref())
166		{
167			*value += public_i * one_minus_alpha;
168		}
169
170		hidden
171	}
172}
173
174/// Computes the phase-2 first-round message: the degree-2 round polynomial that binds the
175/// segment selector, evaluated sparsely without materializing the combined witness.
176///
177/// With `W(X, y) = (1 - X) * P_pad(y) + X * H(y)` and `M` likewise,
178///
179/// ```text
180/// y_1 = sum_y H * M_h    y_inf = sum_y (P_pad + H) * (M_p_pad + M_h)
181/// ```
182///
183/// The dense `H * M_h` pass dominates, and the `y_inf` corrections only have support on the
184/// public prefix.
185/// Both corrections run over whole packed elements, so past the public length the stray terms
186/// the zero padding introduces cancel between the two sums.
187fn first_round_coeffs<F, P: PackedField<Scalar = F>, A: Allocator>(
188	witness: &SegmentPair<'_, P, A>,
189	monster: &SegmentPair<'_, P, A>,
190	gamma: F,
191) -> RoundCoeffs<F>
192where
193	F: BinaryField,
194{
195	// The dense hidden-segment pass.
196	let wide_dense = (witness.hidden.as_ref(), monster.hidden.as_ref())
197		.into_par_iter()
198		.with_min_task(WorkPerItem::FieldMuls)
199		.map(|(&hidden_i, &monster_i)| P::wide_mul(hidden_i, monster_i))
200		.reduce(<P as WideMul>::Output::default, |lhs, rhs| lhs + rhs);
201
202	// The public-prefix corrections.
203	let n_public_packed = witness.public.as_ref().len();
204	let (wide_low_hidden, wide_low_cross) = iter::zip(
205		iter::zip(witness.public.as_ref(), &witness.hidden.as_ref()[..n_public_packed]),
206		iter::zip(monster.public.as_ref(), &monster.hidden.as_ref()[..n_public_packed]),
207	)
208	.map(|((&public_i, &hidden_i), (&public_monster_i, &hidden_monster_i))| {
209		(
210			P::wide_mul(hidden_i, hidden_monster_i),
211			P::wide_mul(public_i + hidden_i, public_monster_i + hidden_monster_i),
212		)
213	})
214	.fold(
215		(<P as WideMul>::Output::default(), <P as WideMul>::Output::default()),
216		|(acc_hidden, acc_cross), (hidden_term, cross_term)| {
217			(acc_hidden + hidden_term, acc_cross + cross_term)
218		},
219	);
220
221	let sum_lanes = |wide: <P as WideMul>::Output| P::reduce(wide).iter().sum::<F>();
222	let y_1 = sum_lanes(wide_dense);
223	let y_inf = y_1 + sum_lanes(wide_low_hidden) + sum_lanes(wide_low_cross);
224
225	RoundEvals([y_1, y_inf]).interpolate(gamma)
226}
227
228/// Executes the phase-2 sumcheck over the witness, with a sparse first round.
229///
230/// # Overview
231///
232/// The witness and the constraint-matrix multilinear are each given as a (public, hidden)
233/// segment pair.
234/// The top word-index variable selects the segment.
235///
236/// The first round binds that selector without materializing the mostly-zero combined buffers.
237/// After the selector challenge, the segment pairs fold into single dense buffers, and a
238/// shared dense-product prover proves the remaining rounds.
239/// So every round message is identical to what a fully dense prover would send.
240///
241/// After the sumcheck, this derives the witness evaluation from the combined evaluation: it
242/// evaluates the public segment directly (cheap, like the verifier does), subtracts its
243/// padded contribution, and scales.
244///
245/// It also divides the three bit-index factors back out of the constraint-matrix evaluation,
246/// leaving the wiring evaluation the verifier's claim is about.
247///
248/// # Returns
249///
250/// The sumcheck's concatenated challenges with the witness evaluation, and the wiring
251/// evaluation for the caller to send.
252#[allow(clippy::too_many_arguments)]
253#[instrument(skip_all, name = "run_sumcheck")]
254pub fn run_sumcheck<F, P: PackedField<Scalar = F>, Channel: IPProverChannel<F>, A: Allocator>(
255	public_folded: &FieldVec<P, A>,
256	hidden_folded: FieldVec<P, A>,
257	public_monster: &FieldVec<P, A>,
258	hidden_monster: FieldVec<P, A>,
259	shift_ind_eval: F,
260	public_words: &[Word],
261	r_j: Vec<F>,
262	gamma: F,
263	channel: &mut Channel,
264	alloc: &A,
265) -> ShiftOutput<F>
266where
267	F: BinaryField,
268{
269	// The hidden pair is the dense one every round iterates over.
270	// So it spans the whole word-index space, and the public pair sits at its base.
271	let log_hidden = hidden_folded.log_len();
272	assert_eq!(hidden_monster.log_len(), log_hidden);
273	assert_eq!(public_monster.log_len(), public_folded.log_len());
274	assert!(public_folded.log_len() <= log_hidden);
275
276	let witness = SegmentPair::<'_, P, A>::new(public_folded, hidden_folded);
277	let monster = SegmentPair::<'_, P, A>::new(public_monster, hidden_monster);
278
279	// Round 1: bind the segment selector.
280	let round_coeffs = first_round_coeffs(&witness, &monster, gamma);
281	channel.send_many(round_coeffs.clone().truncate().coeffs());
282	let alpha = channel.sample();
283	let round_sum = round_coeffs.evaluate(&alpha);
284
285	// Fold the segment pairs at the selector challenge and run the remaining rounds with the
286	// standard prover.
287	let folded_witness = witness.fold(alpha);
288	let folded_monster = monster.fold(alpha);
289	let prover = bivariate_product_prover(alloc, [folded_witness, folded_monster], round_sum);
290
291	let ProveSingleOutput {
292		multilinear_evals,
293		challenges,
294	} = prove_single(prover, channel);
295
296	let mut r_y = iter::once(alpha).chain(challenges).collect::<Vec<_>>();
297	// Reverse the challenges to get the evaluation point.
298	r_y.reverse();
299
300	let [trace_eval, monster_eval] = multilinear_evals
301		.try_into()
302		.expect("prover has 2 multilinear polynomials");
303
304	// Every constraint-matrix entry carries the three bit-index factors.
305	// Dividing them out leaves the bare wiring evaluation.
306	//
307	// Like the witness evaluation below, this makes the protocol incomplete with negligible
308	// probability, when the scale is zero.
309	let wiring_eval = monster_eval * shift_ind_eval.invert_or_zero();
310
311	// Derive the witness evaluation from the combined evaluation: evaluate the public segment
312	// directly (cheap, like the verifier does), subtract its padded contribution, and scale.
313	//
314	// This makes the protocol incomplete with negligible probability, when the segment
315	// selector challenge is zero.
316	let log_half = r_y.len() - 1;
317	let r_segment = r_y[log_half];
318	// Round the public word count up to a power of two: the segment spans that many word
319	// slots.
320	// The count itself need not be a power of two, and the missing words are read as zero.
321	let log_public_words = log2_ceil_usize(public_words.len());
322	let public_eval = evaluate_words_mle::<F, F>(public_words, &r_j, &r_y[..log_public_words]);
323	let padded_public_eval = eq_ind_zero(&r_y[log_public_words..log_half]) * public_eval;
324	let witness_eval =
325		(trace_eval - (F::ONE - r_segment) * padded_public_eval) * r_segment.invert_or_zero();
326	channel.send_one(witness_eval);
327
328	ShiftOutput {
329		sumcheck: SumcheckOutput {
330			challenges: [r_j, r_y].concat(),
331			eval: witness_eval,
332		},
333		wiring_eval,
334	}
335}
336
337/// What the shift reduction leaves for its caller.
338///
339/// The wiring evaluation is not sent here: the verifier reads it after the public segment's
340/// evaluation claim, which the caller proves, so the caller sends it at that point.
341#[derive(Debug)]
342pub struct ShiftOutput<F> {
343	/// The sumcheck's challenges `[r_j, r_y]` and the witness evaluation.
344	pub sumcheck: SumcheckOutput<F>,
345	/// The wiring multilinear's evaluation at the reduced point.
346	pub wiring_eval: F,
347}