Skip to main content

binius_prover/protocols/shift/
phase_1.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{iter, ops::Deref};
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, PackedField, WideMul};
9use binius_ip::sumcheck::RoundCoeffs;
10use binius_ip_prover::{
11	channel::IPProverChannel,
12	sumcheck::{
13		ProveSingleOutput, bivariate_product_evaluator::BivariateProductEvaluator,
14		bivariate_product_prover, common::SumcheckProver, prove_single, round_evals::RoundEvals,
15		round_evaluator::SharedSumcheckProver,
16	},
17};
18use binius_math::{FieldBuffer, FieldVec, multilinear::fold::fold_highest_var_inplace};
19use binius_verifier::protocols::shift::LOG_SHIFT_COUNT;
20use tracing::instrument;
21
22use super::{
23	SegmentWords,
24	claims::PreparedOperandClaims,
25	key_collection::{DenseShiftEncoding, KeyCollection},
26	monster::shift_operator_table,
27	outer::OuterShiftStage,
28	shift_ind::ShiftChallenge,
29};
30
31/// Proves the first phase of the shift reduction.
32///
33/// Builds the witness-and-batching multilinear for both segments and concatenates their rows.
34/// One sumcheck then runs over their product, against a weight table that is never materialized.
35///
36/// # Arguments
37///
38/// - `key_collection`: the prover's key collection for the constraint system.
39/// - `words`: the value-vector words.
40/// - `prepared`: the prepared claim of each operation, indexed by the operation a key names.
41/// - `oblong_weights`: the weights of the reduction's first factor, one per bit position.
42/// - `channel`: the prover channel the interactive rounds run over.
43/// - `alloc`: the allocator the intermediate buffers are drawn from.
44///
45/// # Returns
46///
47/// The challenge point split into its axes, alongside the leftover weights.
48/// Also the two evaluations this phase reduced to.
49#[instrument(skip_all, name = "prover_phase_1")]
50pub fn prove_phase_1<F, P, Channel, A>(
51	key_collection: &KeyCollection,
52	words: SegmentWords<'_>,
53	prepared: &PreparedOperandClaims<F>,
54	oblong_weights: &[F],
55	channel: &mut Channel,
56	alloc: &A,
57) -> Phase1Output<F>
58where
59	F: BinaryField,
60	P: PackedField<Scalar = F>,
61	Channel: IPProverChannel<F>,
62	A: Allocator,
63{
64	// Accumulate the witness-and-batching rows of the public and hidden segments separately.
65	// The public words are the prefix of the value vector.
66	// Each segment's key ranges are relative to its own segment.
67	let public = key_collection
68		.public
69		.build_g::<_, P>(words.public, prepared);
70	let hidden = key_collection
71		.hidden
72		.build_g::<_, P>(words.hidden, prepared);
73	let g = SparseShiftRows::from_segments([
74		(&public, &key_collection.public.dense_shift_enc),
75		(&hidden, &key_collection.hidden.dense_shift_enc),
76	]);
77
78	g.run_phase_1_sumcheck(oblong_weights, prepared.batched_eval, channel, alloc)
79}
80
81/// The number of variables the shift-and-bit phases of the reduction span: the bit position
82/// within a word, the inner shift slot, and the outer shift slot.
83///
84/// No table the prover ever builds spans all three axes at once — avoiding that table is why
85/// the outer-slot rounds exist.
86pub const PHASE_1_LOG_LEN: usize = Word::LOG_BITS + LOG_SHIFT_ROWS;
87
88/// The number of variables of the row-index axis: two shift slots, since a term names two
89/// shifts applied in sequence.
90///
91/// The slots are ordered outer-major, so the rounds binding this axis peel the outer shift
92/// first.
93pub const LOG_SHIFT_ROWS: usize = 2 * LOG_SHIFT_COUNT;
94
95/// The number of variables one shift-weight table spans: one row of weights per (shift
96/// variant, shift amount) pair, one weight per bit position within a word.
97pub const SHIFT_OPERATOR_LOG_LEN: usize = Word::LOG_BITS + LOG_SHIFT_COUNT;
98
99/// The output of the first proving phase.
100///
101/// The challenge point the phase's rounds bound, split into its axes, plus the two
102/// evaluations the reduction still needs to carry forward.
103#[derive(Debug, Clone)]
104pub struct Phase1Output<F> {
105	/// The bit position within a word.
106	pub r_j: Vec<F>,
107	/// The inner shift's amount and variant.
108	pub inner: ShiftChallenge<F>,
109	/// The outer shift's amount and variant.
110	pub outer: ShiftChallenge<F>,
111	/// The weights carried through the outer shift, at the outer challenge point.
112	///
113	/// Read by the next phase's bit-index rounds.
114	pub psi: Vec<F>,
115	/// The evaluation claim the next phase proves: the product of the two evaluations below.
116	pub gamma: F,
117	/// The witness-and-batching multilinear, evaluated at the bound shift and bit point.
118	///
119	/// Carried through the remaining phases' rounds, so it is never recomputed.
120	pub g_eval: F,
121}
122
123/// The number of packed elements one row of `Word::BITS` scalars occupies.
124pub(super) const fn row_len<P: PackedField>() -> usize {
125	assert!(
126		P::LOG_WIDTH <= Word::LOG_BITS,
127		"a row of `Word::BITS` scalars must be a whole number of packed elements"
128	);
129	Word::BITS >> P::LOG_WIDTH
130}
131
132/// The nonzero rows of the witness-and-batching multilinear that a constraint system's shifts
133/// reach: one row of weights per (shift variant, shift amount) pair actually named, each
134/// tagged with its row index.
135///
136/// A constraint system only ever names a few dozen pairs — 16 for a SHA-256 circuit, 40 for
137/// Keccak.
138///
139/// Rows at a repeated index add up wherever the multilinear is read, so two segments' rows can
140/// simply concatenate here with no deduplication step.
141#[derive(Debug, Clone)]
142pub struct SparseShiftRows<P: PackedField> {
143	/// The row index of each stored row, in the same order as the row values below.
144	///
145	/// An index can repeat.
146	indices: Vec<u32>,
147	/// The stored rows end to end, in the same order as the indices above.
148	values: Vec<P>,
149	/// The number of row-index variables the sumcheck has yet to bind.
150	log_rows: usize,
151}
152
153impl<P: PackedField> SparseShiftRows<P> {
154	/// Collects the rows two key segments accumulated, tagged with the shift index each sits at.
155	///
156	/// Each segment's rows arrive in its own dense encoding order, which decodes each position
157	/// back to the shift index this list keys on.
158	/// The two segments' rows simply concatenate: a shift both use appears twice, and the two
159	/// rows add up wherever the multilinear is read later.
160	///
161	/// # Panics
162	///
163	/// Panics unless each segment's row count matches what its encoding accounts for.
164	pub fn from_segments(segments: [(&[P], &DenseShiftEncoding); 2]) -> Self {
165		let mut indices = Vec::new();
166		let mut values = Vec::new();
167
168		// Each segment contributes its own rows, tagged with its own shift indices.
169		for (rows, dense_shift_enc) in segments {
170			assert_eq!(
171				rows.len(),
172				dense_shift_enc.len() * row_len::<P>(),
173				"a segment holds one row per shift its encoding names"
174			);
175			// Decode each stored row's position back to the shift index it belongs to.
176			indices.extend(dense_shift_enc.shift_indices().map(|index| index as u32));
177			// Rows just concatenate: a shift both segments use simply appears twice.
178			values.extend_from_slice(rows);
179		}
180
181		Self::new(indices, values, LOG_SHIFT_ROWS)
182	}
183
184	/// Collects stored rows sitting at the given indices of a row space `log_rows` variables
185	/// wide.
186	///
187	/// An index can repeat: rows at a repeated index add up wherever the multilinear is read
188	/// later.
189	///
190	/// # Panics
191	///
192	/// Panics unless there is one row of values per index and every index fits the row space.
193	pub fn new(indices: Vec<u32>, values: Vec<P>, log_rows: usize) -> Self {
194		assert_eq!(
195			values.len(),
196			indices.len() * row_len::<P>(),
197			"the values hold one row per index"
198		);
199		assert!(
200			indices
201				.iter()
202				.all(|&index| (index as usize) < 1 << log_rows),
203			"every index names a row of the space"
204		);
205
206		Self {
207			indices,
208			values,
209			log_rows,
210		}
211	}
212
213	/// The number of row-index variables the sumcheck has yet to bind.
214	pub(crate) const fn log_rows(&self) -> usize {
215		self.log_rows
216	}
217
218	/// The stored rows, each with the row index it sits at.
219	pub(crate) fn rows(&self) -> impl Iterator<Item = (usize, &[P])> {
220		iter::zip(&self.indices, self.values.chunks_exact(row_len::<P>()))
221			.map(|(&index, row)| (index as usize, row))
222	}
223
224	/// The index-space bit the next round binds, separating the two halves of the row space.
225	///
226	/// # Panics
227	///
228	/// Panics unless at least one row-index variable remains to bind.
229	pub(crate) fn half(&self) -> usize {
230		assert!(self.log_rows > 0, "precondition: a row-index variable remains to bind");
231		1 << (self.log_rows - 1)
232	}
233
234	/// Binds the highest row-index variable to a challenge.
235	///
236	/// Folding is linear, so a row keeps its identity across it: it is scaled by the challenge
237	/// weight of its half, then moved down into the folded index space.
238	/// The list stays the same length, with nothing paired up or merged.
239	///
240	/// # Panics
241	///
242	/// Panics unless at least one row-index variable remains to bind.
243	pub(crate) fn fold(&mut self, challenge: P::Scalar) {
244		let half = self.half();
245		let lower_weight = P::broadcast(P::Scalar::ONE - challenge);
246		let upper_weight = P::broadcast(challenge);
247
248		let row_len = row_len::<P>();
249		// Every stored row moves on its own: the fold is linear, so no row needs its
250		// counterpart from the other half to update.
251		for (position, index) in self.indices.iter_mut().enumerate() {
252			let row = &mut self.values[position * row_len..][..row_len];
253			if *index as usize & half == 0 {
254				// Lower half: scale by `1 - challenge` and keep the same index.
255				row.iter_mut().for_each(|value| *value *= lower_weight);
256			} else {
257				// Upper half: scale by `challenge` and fold the index into the lower half.
258				row.iter_mut().for_each(|value| *value *= upper_weight);
259				*index ^= half as u32;
260			}
261		}
262
263		self.log_rows -= 1;
264	}
265
266	/// Collapses every remaining row-index variable into one dense row over the bit position:
267	/// the sum of the stored rows, now that they all sit at the same index.
268	///
269	/// # Panics
270	///
271	/// Panics unless every row-index variable is already bound.
272	fn into_bit_multilinear<A: Allocator>(self, alloc: &A) -> FieldVec<P, A> {
273		assert_eq!(self.log_rows, 0, "precondition: every row-index variable is bound");
274
275		let mut multilinear = FieldBuffer::zeros_in(alloc, Word::LOG_BITS);
276		// Every stored row now sits at the same index, so they all add into one dense row.
277		for (_, row) in self.rows() {
278			for (slot, &value) in iter::zip(multilinear.as_mut(), row) {
279				*slot += value;
280			}
281		}
282		multilinear
283	}
284
285	/// Computes one round message of the phase-1 sumcheck over the row index.
286	///
287	/// Sampled at 1 and at infinity, with the claim supplying the value at 0.
288	/// Both samples are linear in the stored rows, so each row contributes independently and
289	/// facing rows across the split never have to be paired up.
290	///
291	/// ```text
292	/// R(1)   = sum_v G_1(v) H_1(v)             row (i, c) adds <c, h[i]>, upper half only
293	/// R(inf) = sum_v (G_0 + G_1)(H_0 + H_1)    row (i, c) adds <c, h[i] + h[i ^ half]>, either half
294	/// ```
295	fn round_coeffs<F, Data>(&self, h: &FieldBuffer<P, Data>, claim: F) -> RoundCoeffs<F>
296	where
297		F: Field,
298		P: PackedField<Scalar = F>,
299		Data: Deref<Target = [P]>,
300	{
301		let half = self.half();
302		let row_len = row_len::<P>();
303		let h_rows = h.as_ref();
304
305		// The per-point products accumulate in unreduced (wide) form and reduce once at the
306		// end.
307		let mut y_1 = <P as WideMul>::Output::default();
308		let mut y_inf = <P as WideMul>::Output::default();
309		for (index, row) in self.rows() {
310			let own = &h_rows[index * row_len..][..row_len];
311			let facing = &h_rows[(index ^ half) * row_len..][..row_len];
312
313			for i in 0..row_len {
314				// The infinity evaluation reads H(0) + H(1), the same sum from either half.
315				// So a row's own half only decides its contribution to the evaluation at 1.
316				if index & half != 0 {
317					y_1 += P::wide_mul(row[i], own[i]);
318				}
319				y_inf += P::wide_mul(row[i], own[i] + facing[i]);
320			}
321		}
322
323		// A row is a whole number of packed elements, so every lane of the accumulators is
324		// live.
325		let sum_lanes = |wide| P::reduce(wide).iter().sum::<F>();
326		RoundEvals([sum_lanes(y_1), sum_lanes(y_inf)]).interpolate(claim)
327	}
328
329	/// Runs the phase-1 sumcheck over the product of this row list and a weight table.
330	///
331	/// This row list is zero outside the rows a constraint system's shifts name, dense within
332	/// a named row.
333	/// A weight table holding one row per possible shift sequence would need `2^24` entries
334	/// and is never built:
335	///
336	/// - The rounds binding the outer shift slot read only the stored rows, so their cost follows
337	///   the shifts the constraint system names, not the whole space.
338	/// - Once the outer slot is bound, both multilinears are one dense row, and a shared
339	///   dense-product prover runs the remaining rounds.
340	///
341	/// Every round message matches what a fully dense prover would send.
342	///
343	/// The outer rounds run here, rather than inside that dense-product prover, because the
344	/// weights they leave behind must outlive it: the next phase's rounds run against those
345	/// same weights.
346	///
347	/// # Arguments
348	///
349	/// - `oblong_weights`: the weights of the reduction's first factor, one per bit position.
350	/// - `sum`: the claim being proved, which must equal the true product exactly when the witness
351	///   satisfies the constraint system.
352	///
353	/// # Returns
354	///
355	/// The challenge point split into its axes, the leftover weights, and the two evaluations
356	/// this phase reduced to.
357	#[instrument(skip_all, name = "run_sumcheck")]
358	pub fn run_phase_1_sumcheck<F, Channel, A>(
359		mut self,
360		oblong_weights: &[F],
361		sum: F,
362		channel: &mut Channel,
363		alloc: &A,
364	) -> Phase1Output<F>
365	where
366		F: BinaryField,
367		P: PackedField<Scalar = F>,
368		Channel: IPProverChannel<F>,
369		A: Allocator,
370	{
371		assert_eq!(self.log_rows(), LOG_SHIFT_ROWS, "the row list spans both shift slots");
372
373		// Phase 1a: bind the outer shift slot, one round at a time.
374		//
375		// Each round asks the outer-slot driver for the round polynomial, sends it, samples
376		// the challenge, then folds both the driver and the row list by that challenge.
377		let mut outer = OuterShiftStage::new(alloc, oblong_weights);
378		let mut claim = sum;
379		let mut outer_point = Vec::with_capacity(LOG_SHIFT_COUNT);
380		for _ in 0..LOG_SHIFT_COUNT {
381			let round_coeffs = outer.round_coeffs(&self, claim);
382			channel.send_many(round_coeffs.clone().truncate().coeffs());
383			let challenge = channel.sample();
384			claim = round_coeffs.evaluate(&challenge);
385			outer.fold(challenge);
386			self.fold(challenge);
387			outer_point.push(challenge);
388		}
389		// The weights the outer rounds leave behind: what the next phase runs against.
390		let psi = outer.psi().to_vec();
391
392		// Phase 1b: bind the inner shift slot and the bit position.
393		//
394		// What is left is one shift slot and the bit position, against the weight table built
395		// from the outer rounds' leftover weights.
396		let h = shift_operator_table(alloc, &psi);
397
398		// The row list itself becomes the sparse half of the row-and-bit sumcheck prover.
399		let g = self;
400		let ProveSingleOutput {
401			multilinear_evals,
402			mut challenges,
403		} = prove_single(Phase1SumcheckProver::new(g, h, claim, alloc), channel);
404
405		// The rounds bind coordinates from the most significant one down.
406		// Reversing recovers the evaluation point in increasing order of significance: the bit
407		// position, then the inner slot's amount and variant, then the outer slot's.
408		challenges.reverse();
409		assert_eq!(challenges.len(), SHIFT_OPERATOR_LOG_LEN);
410		let mut r_j = challenges;
411		let r_v_inner = r_j.split_off(Word::LOG_BITS * 2);
412		let r_s_inner = r_j.split_off(Word::LOG_BITS);
413
414		outer_point.reverse();
415		let mut r_s_outer = outer_point;
416		let r_v_outer = r_s_outer.split_off(Word::LOG_BITS);
417
418		let [g_eval, h_eval] = multilinear_evals
419			.try_into()
420			.expect("prover has 2 multilinear polynomials");
421
422		Phase1Output {
423			r_j,
424			inner: ShiftChallenge::new(r_s_inner, r_v_inner),
425			outer: ShiftChallenge::new(r_s_outer, r_v_outer),
426			psi,
427			gamma: g_eval * h_eval,
428			g_eval,
429		}
430	}
431}
432
433/// A sumcheck prover for the product of a sparse row list and a dense weight table, both
434/// spanning one shift slot and the bit position.
435///
436/// The row list is zero outside the rows a constraint system's shifts name, dense within a
437/// named row.
438/// So this prover changes strategy halfway through:
439///
440/// - While the row index has unbound variables, each round reads only the stored rows, at a cost
441///   that follows how many shifts the constraint system names, not the whole row space.
442/// - Once the row index is bound, both multilinears are one dense row, and a shared dense-product
443///   prover takes over.
444///
445/// Every round message matches what a fully dense prover would send.
446///
447/// # Performance
448///
449/// Both the sampled evaluations and the fold are linear in the sparse row list, so each
450/// stored row contributes independently, walking whole packed rows rather than individual
451/// scalar points.
452pub struct Phase1SumcheckProver<'alloc, A: Allocator, P: PackedField> {
453	alloc: &'alloc A,
454	/// The stage the protocol is in.
455	///
456	/// Always holds a value between calls, only briefly emptied while the row-stage buffers
457	/// move into the bit-stage prover.
458	stage: Option<Stage<'alloc, A, P>>,
459}
460
461/// Which half of the protocol the prover is in.
462enum Stage<'alloc, A: Allocator, P: PackedField> {
463	/// Binding the row index, over the rows the sparse list stores.
464	Rows {
465		g: SparseShiftRows<P>,
466		h: FieldVec<P, A>,
467		/// This round's sum claim.
468		claim: P::Scalar,
469		/// The round polynomial.
470		///
471		/// Set once the round's message is produced, and cleared again once the challenge
472		/// that reduces it to the next claim arrives.
473		coeffs: Option<RoundCoeffs<P::Scalar>>,
474	},
475	/// Binding the bit position, with both multilinears now one dense row.
476	Bits(SharedSumcheckProver<'alloc, A, P, BivariateProductEvaluator>),
477}
478
479impl<'alloc, A: Allocator, F: Field, P: PackedField<Scalar = F>>
480	Phase1SumcheckProver<'alloc, A, P>
481{
482	/// Creates a prover for the claim that the two multilinears' product sums to the given
483	/// value over the hypercube.
484	///
485	/// The outer shift slot is already bound by the time this runs, so the dense multilinear
486	/// is the weight table over the weights those rounds left behind, and the row list's
487	/// remaining row index is the inner slot.
488	///
489	/// # Panics
490	///
491	/// Panics unless both multilinears span exactly one shift slot and the bit position.
492	pub fn new(g: SparseShiftRows<P>, h: FieldVec<P, A>, sum: F, alloc: &'alloc A) -> Self {
493		assert_eq!(h.log_len(), SHIFT_OPERATOR_LOG_LEN, "h spans one shift slot");
494		assert_eq!(g.log_rows(), LOG_SHIFT_COUNT, "g's rows span one shift slot");
495
496		Self {
497			alloc,
498			stage: Some(Stage::Rows {
499				g,
500				h,
501				claim: sum,
502				coeffs: None,
503			}),
504		}
505	}
506}
507
508impl<A: Allocator, F: Field, P: PackedField<Scalar = F>> SumcheckProver<F>
509	for Phase1SumcheckProver<'_, A, P>
510{
511	fn n_vars(&self) -> usize {
512		match self.stage.as_ref().expect("the stage is set between calls") {
513			// The weight table spans both axes through the row rounds.
514			// So its length is what remains.
515			Stage::Rows { h, .. } => h.log_len(),
516			Stage::Bits(prover) => prover.n_vars(),
517		}
518	}
519
520	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
521		match self.stage.as_mut().expect("the stage is set between calls") {
522			Stage::Rows {
523				g,
524				h,
525				claim,
526				coeffs,
527			} => {
528				// Row stage: compute this round's message from the stored rows alone, and
529				// hold onto it until the challenge that reduces the claim arrives.
530				let round_coeffs = g.round_coeffs(h, *claim);
531				*coeffs = Some(round_coeffs.clone());
532				vec![round_coeffs]
533			}
534			// Bit stage: delegate to the shared dense-product prover.
535			Stage::Bits(prover) => prover.execute(),
536		}
537	}
538
539	fn fold(&mut self, challenge: F) {
540		// Taken out so the row stage's buffers can move into the bit stage's prover below.
541		let stage = self.stage.take().expect("the stage is set between calls");
542		self.stage = Some(match stage {
543			Stage::Rows {
544				mut g,
545				mut h,
546				coeffs,
547				..
548			} => {
549				let claim = coeffs
550					.expect("execute is called before fold")
551					.evaluate(&challenge);
552				// Fold both multilinears by the same challenge.
553				// The sparse rows move into their half's slot of the shrunken row space.
554				// The dense weight table folds its highest variable the ordinary way.
555				g.fold(challenge);
556				fold_highest_var_inplace(&mut h, challenge);
557
558				if g.log_rows() > 0 {
559					// The row index still has unbound variables: stay in the row stage.
560					Stage::Rows {
561						g,
562						h,
563						claim,
564						coeffs: None,
565					}
566				} else {
567					// The row index is bound.
568					// What is left of each multilinear is one dense row over the bit
569					// position, which the shared prover handles from here.
570					Stage::Bits(bivariate_product_prover(
571						self.alloc,
572						[g.into_bit_multilinear(self.alloc), h],
573						claim,
574					))
575				}
576			}
577			Stage::Bits(mut prover) => {
578				// Bit stage: delegate the fold to the shared dense-product prover.
579				prover.fold(challenge);
580				Stage::Bits(prover)
581			}
582		});
583	}
584
585	fn finish(self) -> Vec<F> {
586		match self.stage.expect("the stage is set between calls") {
587			Stage::Rows { .. } => panic!("finish called before the row index was bound"),
588			// The columns went in as `[g, h]`, so the evaluations come out in that order.
589			Stage::Bits(prover) => prover.finish(),
590		}
591	}
592}
593
594#[cfg(test)]
595mod tests {
596	use binius_compute::GlobalAllocator;
597	use binius_core::constraint_system::{
598		AndConstraint, ConstraintSystem, InoutSegment, Shift, ShiftedValueIndex, ValueIndex,
599	};
600	use binius_field::{Field, Ghash128b, PackedGhash2x128b};
601	use binius_math::{inner_product::inner_product_buffers, test_utils::random_scalars};
602	use binius_transcript::ProverTranscript;
603	use binius_verifier::config::StdChallenger;
604	use rand::{SeedableRng, rngs::StdRng};
605
606	use super::*;
607	use crate::protocols::shift::KeyCollection;
608
609	type F = Ghash128b;
610
611	impl<P: PackedField> SparseShiftRows<P> {
612		/// Spreads the rows over the space they still span, as a dense multilinear.
613		///
614		/// Rows at a repeated index add up.
615		/// Every row the constraint system does not name stays zero.
616		///
617		/// Nothing in the actual proving path needs this.
618		/// It exists only so these tests have a dense reference to check the sparse rounds
619		/// against.
620		fn scatter<A: Allocator>(&self, alloc: &A) -> FieldVec<P, A> {
621			let row_len = row_len::<P>();
622			let mut g = FieldBuffer::zeros_in(alloc, self.log_rows + Word::LOG_BITS);
623			for (index, row) in self.rows() {
624				// A row is a whole number of packed elements, so it lands at row alignment.
625				let slots = &mut g.as_mut()[index * row_len..][..row_len];
626				for (slot, &value) in iter::zip(slots, row) {
627					*slot += value;
628				}
629			}
630			g
631		}
632	}
633
634	/// A system whose two segments name overlapping but distinct shifts.
635	///
636	/// The public segment names `(Sll, 0)` and `(Slr, 3)`; the hidden one `(Sll, 0)`, `(Sar, 7)`
637	/// and `(Rotr, 1)`. So `(Sll, 0)` is the shift the merge has to sum across segments.
638	fn overlapping_shift_system() -> ConstraintSystem {
639		let public = ValueIndex::constant(1);
640		let hidden = ValueIndex::private(1);
641		ConstraintSystem {
642			constants: vec![Word::ZERO; 4],
643			n_inout: 0,
644			n_private: 4,
645			zero_constraints: Vec::new(),
646			and_constraints: vec![AndConstraint([
647				vec![
648					ShiftedValueIndex::plain(public),
649					ShiftedValueIndex::srl(public, 3),
650				],
651				vec![ShiftedValueIndex::sar(hidden, 7)],
652				vec![
653					ShiftedValueIndex::rotr(hidden, 1),
654					ShiftedValueIndex::plain(hidden),
655				],
656			])],
657			imul_constraints: Vec::new(),
658			bmul_constraints: Vec::new(),
659		}
660	}
661
662	/// The two segments' rows concatenate, each tagged with the shift index it sits at.
663	///
664	/// A shift both segments name appears twice rather than being merged — `g` is the sum of its
665	/// rows, so the two add up wherever it is read, and nothing has to deduplicate them.
666	#[test]
667	fn from_segments_concatenates_the_two_encodings() {
668		let key_collection =
669			KeyCollection::build(&overlapping_shift_system(), InoutSegment::Public);
670
671		// Fill each segment's rows with a distinct constant per row, so a row's value says which
672		// segment it came from.
673		let segment_rows = |enc: &DenseShiftEncoding, base: u128| {
674			(0..enc.len() * Word::BITS)
675				.map(|i| F::new(base + (i / Word::BITS) as u128))
676				.collect::<Vec<F>>()
677		};
678		let public = segment_rows(&key_collection.public.dense_shift_enc, 0x100);
679		let hidden = segment_rows(&key_collection.hidden.dense_shift_enc, 0x200);
680
681		let g = SparseShiftRows::from_segments([
682			(&public, &key_collection.public.dense_shift_enc),
683			(&hidden, &key_collection.hidden.dense_shift_enc),
684		]);
685
686		// The public segment's two shifts, then the hidden segment's three. Every term here is
687		// singly shifted, so its outer slot is the identity and its quadruple index is its inner
688		// shift's. `(Sll, 0)` is the first row of both segments, so index 0 appears twice.
689		let row_index = |shift: Shift| shift.index() as u32;
690		assert_eq!(
691			g.indices,
692			[
693				row_index(Shift::IDENTITY),
694				row_index(Shift::srl(3)),
695				row_index(Shift::IDENTITY),
696				row_index(Shift::sar(7)),
697				row_index(Shift::rotr(1)),
698			]
699		);
700
701		// Each row is the one its own segment accumulated, untouched by the other's.
702		let row = |position: usize| &g.values[position * Word::BITS..][..Word::BITS];
703		for (position, expected) in [0x100, 0x101, 0x200, 0x201, 0x202].into_iter().enumerate() {
704			assert!(row(position).iter().all(|&value| value == F::new(expected)));
705		}
706
707		// Where `g` is read, the two rows at the identity add up.
708		let at = |shift: Shift| {
709			g.rows()
710				.filter(|&(index, _)| index == shift.index())
711				.map(|(_, row)| row[0])
712				.sum::<F>()
713		};
714		assert_eq!(at(Shift::IDENTITY), F::new(0x100) + F::new(0x200));
715		assert_eq!(at(Shift::srl(3)), F::new(0x101));
716		assert_eq!(at(Shift::sar(7)), F::new(0x201));
717		assert_eq!(at(Shift::rotr(1)), F::new(0x202));
718	}
719
720	/// A sequence is placed outer-major, so the outer slot lands where the first rounds bind it.
721	#[test]
722	fn a_sequence_is_keyed_outer_major() {
723		let hidden = ValueIndex::private(1);
724		let sequence = [Shift::srl(3), Shift::sll(5)];
725		let cs = ConstraintSystem {
726			constants: vec![Word::ZERO; 4],
727			n_inout: 0,
728			n_private: 4,
729			zero_constraints: Vec::new(),
730			and_constraints: vec![AndConstraint([
731				vec![ShiftedValueIndex::new(hidden, sequence)],
732				Vec::new(),
733				Vec::new(),
734			])],
735			imul_constraints: Vec::new(),
736			bmul_constraints: Vec::new(),
737		};
738
739		let key_collection = KeyCollection::build(&cs, InoutSegment::Public);
740		let [inner, outer] = sequence;
741		assert_eq!(
742			key_collection
743				.hidden
744				.dense_shift_enc
745				.shift_indices()
746				.collect::<Vec<_>>(),
747			[outer.index() << LOG_SHIFT_COUNT | inner.index()]
748		);
749	}
750
751	/// The scatter puts every row where its shift index names, and leaves the rest zero.
752	///
753	/// The outer rounds are past by the time the dense reference below is taken, so the rows span
754	/// one shift slot here rather than the quadruple `from_segments` keys on.
755	#[test]
756	fn scatter_places_rows_at_their_shift_index() {
757		let indices = [Shift::IDENTITY, Shift::sar(7), Shift::rotr(1)]
758			.map(|shift| shift.index() as u32)
759			.to_vec();
760		let values = (0..indices.len() * Word::BITS)
761			.map(|i| F::new(1 + (i / Word::BITS) as u128))
762			.collect::<Vec<F>>();
763		let rows = SparseShiftRows::<F>::new(indices.clone(), values, LOG_SHIFT_COUNT);
764
765		let g = rows.scatter(&GlobalAllocator);
766		assert_eq!(g.log_len(), SHIFT_OPERATOR_LOG_LEN);
767
768		// Exactly the named rows are non-zero, and each holds what the list held.
769		for (row, &shift_index) in indices.iter().enumerate() {
770			let offset = shift_index as usize * Word::BITS;
771			for bit in 0..Word::BITS {
772				assert_eq!(g.get(offset + bit), F::new(1 + row as u128));
773			}
774		}
775		for row in (0..1 << LOG_SHIFT_COUNT).filter(|row| !indices.contains(&(*row as u32))) {
776			assert!((0..Word::BITS).all(|bit| g.get(row * Word::BITS + bit) == F::ZERO));
777		}
778	}
779
780	/// Runs the same claim through a single dense product-sumcheck prover, with no round handled
781	/// specially.
782	///
783	/// This is the reference the sparse-and-dense hybrid prover's round messages have to
784	/// reproduce.
785	fn run_dense_reference<P: PackedField<Scalar = F>>(
786		g: &SparseShiftRows<P>,
787		h: FieldVec<P, GlobalAllocator>,
788		sum: F,
789		channel: &mut ProverTranscript<StdChallenger>,
790	) -> (Vec<F>, F) {
791		let prover =
792			bivariate_product_prover(&GlobalAllocator, [g.scatter(&GlobalAllocator), h], sum);
793
794		let ProveSingleOutput {
795			multilinear_evals,
796			mut challenges,
797		} = prove_single(prover, channel);
798		challenges.reverse();
799
800		let [g_eval, h_eval] = multilinear_evals
801			.try_into()
802			.expect("prover has 2 multilinear polynomials");
803
804		(challenges, g_eval * h_eval)
805	}
806
807	/// The `g` and `h` the row and bit rounds run over, from pseudo-random weights.
808	///
809	/// The outer rounds are past by this point, so `g`'s rows sit at the inner shift each sequence
810	/// names. They carry arbitrary values rather than ones a witness produces: the sumcheck is a
811	/// statement about whatever multilinears it is handed, so what the rows hold does not bear on
812	/// whether the sparse rounds reproduce the dense ones.
813	fn phase_1_multilinears<P: PackedField<Scalar = F>>(
814		cs: &ConstraintSystem,
815		seed: u64,
816	) -> (SparseShiftRows<P>, FieldVec<P, GlobalAllocator>) {
817		let mut rng = StdRng::seed_from_u64(seed);
818		let key_collection = KeyCollection::build(cs, InoutSegment::Public);
819
820		let mut indices = Vec::new();
821		let mut values = Vec::new();
822		for segment in [&key_collection.public, &key_collection.hidden] {
823			for [inner, _] in segment.dense_shift_enc.iter() {
824				indices.push(inner.index() as u32);
825				values.extend((0..row_len::<P>()).map(|_| P::random(&mut rng)));
826			}
827		}
828		let g = SparseShiftRows::new(indices, values, LOG_SHIFT_COUNT);
829
830		// The weights `h` is built from are arbitrary here for the same reason `g`'s rows are.
831		let h = shift_operator_table(&GlobalAllocator, &random_scalars::<F>(&mut rng, Word::BITS));
832
833		(g, h)
834	}
835
836	/// The sparse rounds send exactly what the dense prover sends.
837	///
838	/// The two provers sum the same product over the same hypercube, so every round message — and
839	/// therefore the whole transcript, the challenges it draws, and the evaluation it reduces to —
840	/// must agree. This is what lets the verifier stay untouched.
841	fn assert_sparse_matches_dense<P: PackedField<Scalar = F>>(cs: &ConstraintSystem, seed: u64) {
842		let (g, h) = phase_1_multilinears::<P>(cs, seed);
843		// The true sum, so the test exercises a sumcheck a verifier would accept.
844		let sum = inner_product_buffers(&g.scatter(&GlobalAllocator), &h);
845
846		let mut sparse_transcript = ProverTranscript::<StdChallenger>::default();
847		let ProveSingleOutput {
848			multilinear_evals,
849			challenges: mut sparse_challenges,
850		} = prove_single(
851			Phase1SumcheckProver::new(g.clone(), h.clone(), sum, &GlobalAllocator),
852			&mut sparse_transcript,
853		);
854		sparse_challenges.reverse();
855		let [g_eval, h_eval] = multilinear_evals
856			.try_into()
857			.expect("prover has 2 multilinear polynomials");
858
859		let mut dense_transcript = ProverTranscript::<StdChallenger>::default();
860		let (dense_challenges, dense_eval) = run_dense_reference(&g, h, sum, &mut dense_transcript);
861
862		assert_eq!(sparse_challenges, dense_challenges);
863		assert_eq!(g_eval * h_eval, dense_eval);
864		assert_eq!(sparse_transcript.finalize(), dense_transcript.finalize());
865	}
866
867	#[test]
868	fn sparse_rounds_match_the_dense_prover() {
869		// Both packing widths the prover is instantiated at: the m4 prover drives phase 1 with
870		// scalars, the single-instance prover with a packed field.
871		assert_sparse_matches_dense::<F>(&overlapping_shift_system(), 0);
872		assert_sparse_matches_dense::<PackedGhash2x128b>(&overlapping_shift_system(), 1);
873	}
874
875	/// A constraint system that constrains nothing names no shift, so `g` stores no row at all.
876	///
877	/// The sparse rounds then have nothing to scan, which must still leave the transcript the one
878	/// a dense prover over the zero `g` would have written.
879	#[test]
880	fn sparse_rounds_match_the_dense_prover_with_an_empty_g() {
881		let cs = ConstraintSystem {
882			constants: vec![Word::ZERO; 4],
883			n_inout: 0,
884			n_private: 4,
885			zero_constraints: Vec::new(),
886			and_constraints: Vec::new(),
887			imul_constraints: Vec::new(),
888			bmul_constraints: Vec::new(),
889		};
890
891		let key_collection = KeyCollection::build(&cs, InoutSegment::Public);
892		assert!(key_collection.public.dense_shift_enc.is_empty());
893		assert!(key_collection.hidden.dense_shift_enc.is_empty());
894
895		assert_sparse_matches_dense::<PackedGhash2x128b>(&cs, 2);
896	}
897}