Skip to main content

binius_prover/protocols/shift/key_collection/
collection.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use binius_compute::{Allocator, VecLike};
5use binius_core::constraint_system::{ConstraintSystem, InoutSegment};
6use binius_field::{BinaryField, PackedField, WideMul};
7use binius_math::{FieldBuffer, FieldVec, multilinear::eq::eq_ind_partial_eval};
8use binius_utils::{
9	checked_arithmetics::log2_ceil_usize,
10	rayon::{
11		prelude::*,
12		task_size::{IndexedParallelIteratorExt, WorkPerItem},
13	},
14	serialization::{DeserializeBytes, SerializationError, SerializeBytes},
15};
16use bytes::{Buf, BufMut};
17use tracing::instrument;
18
19use super::{builder, dense_shift_encoding::DenseShiftEncoding, key_segment::KeySegment};
20use crate::protocols::shift::{
21	claims::PreparedOperandClaims, monster::OuterSlotWeights, shift_ind::ShiftChallenge,
22};
23
24/// The prover's complete view of a constraint system's shift keys, split by value-vector segment.
25///
26/// - The public segment covers value-vector indices from 0 up to the public word count.
27/// - The hidden segment covers the rest, up to the combined length.
28/// - Word indices inside each segment are relative to that segment's own start.
29/// - Both phases of the shift reduction iterate the two segments in absolute value-vector order.
30#[derive(Debug, Clone)]
31pub struct KeyCollection {
32	/// The keys of the public segment: constants and inout words.
33	pub public: KeySegment,
34	/// The keys of the hidden segment: private words.
35	pub hidden: KeySegment,
36}
37
38impl KeyCollection {
39	/// Walks a constraint system, collecting every shift key into its segment.
40	///
41	/// Runs as a stable counting sort by word: one pass counts the references each word
42	/// receives, one pass scatters them into a flat array in walk order, and the per-word
43	/// grouping into keys runs over disjoint word ranges in parallel.
44	///
45	/// # Arguments
46	///
47	/// - `cs`: the constraint system to walk.
48	/// - `inout`: the split point between the public and hidden segments.
49	pub fn build(cs: &ConstraintSystem, inout: InoutSegment) -> Self {
50		builder::build_key_collection(cs, inout)
51	}
52
53	/// The base-2 logarithm of the hidden segment length in words, rounded up to a power of two.
54	///
55	/// Matches the corresponding quantity for the constraint system the collection was built from.
56	/// That system guarantees this is at least the public segment's logarithm.
57	///
58	/// ```text
59	/// log_witness_words = ceil_log2( hidden segment length in words )
60	/// ```
61	pub const fn log_witness_words(&self) -> usize {
62		log2_ceil_usize(self.hidden.n_words())
63	}
64
65	/// Builds the constraint-matrix multilinear's two segments.
66	///
67	/// For each witness word, this sums every one of its keys' contributions.
68	/// A key's contribution is its constraint-index accumulation, scaled by a scalar that
69	/// folds together the operand batching weight, the bit-index evaluation, and the two
70	/// equality-indicator weights of the shift sequence the key names.
71	///
72	/// # Returns
73	///
74	/// The public segment, spanning the rounded-up public word count, and the hidden segment,
75	/// spanning the full witness word count.
76	/// The phase-2 sumcheck's sparse first round consumes both directly, without
77	/// materializing their combined buffer.
78	#[instrument(skip_all, name = "build_monster_segments")]
79	pub fn build_monster_segments<F, P: PackedField<Scalar = F>, A: Allocator>(
80		&self,
81		alloc: &A,
82		prepared: &PreparedOperandClaims<F>,
83		h_eval: F,
84		inner: &ShiftChallenge<F>,
85		outer: &ShiftChallenge<F>,
86	) -> (FieldVec<P, A>, FieldVec<P, A>)
87	where
88		F: BinaryField,
89	{
90		let r_v_tensor = eq_ind_partial_eval::<F>(&inner.variant);
91		let r_s_tensor = eq_ind_partial_eval::<F>(&inner.amount);
92
93		// Invariant: a key's sequence weight factorizes across its two slots.
94		//
95		//     eq(r_v1, v_1) * eq(r_s1, s_1)  *  eq(r_v2, v_2) * eq(r_s2, s_2)
96		//     \______ inner slot _________/     \______ outer slot ________/
97		//
98		// So one table per slot suffices, at `2 * SHIFT_COUNT` entries instead of
99		// `SHIFT_COUNT^2`.
100		let outer_weights = OuterSlotWeights::<F>::new(outer);
101
102		// The scalars of one key segment, laid out with the operand column innermost, so a key's
103		// weights form one contiguous chunk its wide accumulation can index by column.
104		//
105		// A key's sequence selects itself through an equality indicator over both slots.
106		// The h evaluation is one factor shared by every key.
107		let operand_stride = prepared.operand_weights.len();
108		let build_scalars = |dense_shift_enc: &DenseShiftEncoding| {
109			dense_shift_enc
110				.iter()
111				.flat_map(|[inner_shift, outer_shift]| {
112					let shift_scalar = h_eval
113						* r_v_tensor.as_ref()[inner_shift.variant as usize]
114						* r_s_tensor.as_ref()[inner_shift.amount as usize]
115						* outer_weights.weight(outer_shift);
116					prepared
117						.operand_weights
118						.iter()
119						.map(move |operand_weight| *operand_weight * shift_scalar)
120				})
121				.collect::<Vec<_>>()
122		};
123
124		// The scalar for one word of a segment: the accumulated contribution of all its
125		// keys, summed unreduced and reduced once at the end.
126		let word_scalar = |segment: &KeySegment, scalars: &[F], index: usize| {
127			let wide = segment
128				.word_keys(index)
129				.iter()
130				.map(|key| {
131					// One chunk per shift sequence, at the padded stride; the slots past the last
132					// column name no operand and are never indexed.
133					let base = key.dense_shift_idx as usize * operand_stride;
134					key.accumulate_wide(
135						&segment.constraint_indices,
136						&prepared.r_x_tensor,
137						&scalars[base..base + operand_stride],
138					)
139				})
140				.sum::<<F as WideMul>::Output>();
141			F::reduce(wide)
142		};
143
144		// Each segment sits at the base of its buffer: the public piece fills its
145		// power-of-two length exactly, the hidden piece is zero-padded up to the hidden
146		// segment length.
147		let build_segment = |segment: &KeySegment, log_len: usize| {
148			// Each segment has its own dense shift encoding, so it has its own scalar table.
149			let scalars = build_scalars(&segment.dense_shift_enc);
150			let capacity = 1 << log_len.saturating_sub(P::LOG_WIDTH);
151			let n_words = segment.n_words();
152			// Full packed elements: each maps exactly `P::WIDTH` words, so `from_scalars`
153			// sees a statically-sized iterator.
154			// The trailing partial element is filled separately below.
155			let n_full = n_words / P::WIDTH;
156			// Allocate the buffer once, then fill its aligned packed elements in parallel
157			// through its spare capacity.
158			let mut values = alloc.alloc::<P>(capacity);
159			values.spare_capacity_mut()[..n_full]
160				.par_iter_mut()
161				.enumerate()
162				.with_min_task(WorkPerItem::FieldMuls)
163				.for_each(|(chunk_index, slot)| {
164					let start = chunk_index * P::WIDTH;
165					slot.write(P::from_scalars(
166						(0..P::WIDTH).map(|i| word_scalar(segment, &scalars, start + i)),
167					));
168				});
169			// Safety: the parallel loop above initialized every one of the `n_full`
170			// slots.
171			unsafe { values.set_len(n_full) };
172			if !n_words.is_multiple_of(P::WIDTH) {
173				let start = n_full * P::WIDTH;
174				values.push(P::from_scalars(
175					(start..n_words).map(|word_index| word_scalar(segment, &scalars, word_index)),
176				));
177			}
178			values.resize(capacity, P::default());
179			FieldBuffer::new(log_len, values)
180		};
181
182		// The segment word count need not be a power of two; the constraint-matrix
183		// multilinear spans the rounded-up count.
184		let log_public_words = log2_ceil_usize(self.public.n_words());
185		let public_monster = build_segment(&self.public, log_public_words);
186		let hidden_monster = build_segment(&self.hidden, self.log_witness_words());
187
188		(public_monster, hidden_monster)
189	}
190}
191
192impl SerializeBytes for KeyCollection {
193	fn serialize(&self, mut write_buf: impl BufMut) -> Result<(), SerializationError> {
194		// Version for forward compatibility; version 3 introduced the dense shift encoding, and
195		// version 4 dropped the operation from keys, which now span every operand column.
196		const VERSION: u32 = 4;
197		VERSION.serialize(&mut write_buf)?;
198
199		self.public.serialize(&mut write_buf)?;
200		self.hidden.serialize(write_buf)
201	}
202}
203
204impl DeserializeBytes for KeyCollection {
205	fn deserialize(mut read_buf: impl Buf) -> Result<Self, SerializationError> {
206		const VERSION: u32 = 4;
207		let version = u32::deserialize(&mut read_buf)?;
208		if version != VERSION {
209			return Err(SerializationError::InvalidConstruction {
210				name: "KeyCollection::version",
211			});
212		}
213
214		let public = KeySegment::deserialize(&mut read_buf)?;
215		let hidden = KeySegment::deserialize(read_buf)?;
216
217		Ok(KeyCollection { public, hidden })
218	}
219}
220
221#[cfg(test)]
222mod tests {
223	use binius_core::{
224		constraint_system::{AndConstraint, Shift, ShiftedValueIndex, ValueIndex},
225		word::Word,
226	};
227	use binius_verifier::protocols::shift::SHIFT_COUNT;
228
229	use super::*;
230
231	/// A shift sequence carrying one shift, which the canonical form places in the inner slot.
232	fn single(shift: Shift) -> [Shift; 2] {
233		[shift, Shift::IDENTITY]
234	}
235
236	/// A constraint system with a handful of distinct shifts, differing between the two segments.
237	///
238	/// The public segment references `Sll(0)` and `Slr(3)`.
239	/// The hidden one references `Sll(0)`, `Sar(7)` and `Rotr(1)`.
240	/// Every outer slot is the identity, so the sequences sort by their inner shift alone.
241	fn shifted_constraint_system() -> ConstraintSystem {
242		// The system has four constants and no inout values, so the public segment is the
243		// constants and the hidden one is the private values.
244		let public = ValueIndex::constant(1);
245		let hidden = ValueIndex::private(1);
246		ConstraintSystem {
247			constants: vec![Word::ZERO; 4],
248			n_inout: 0,
249			n_private: 4,
250			zero_constraints: Vec::new(),
251			and_constraints: vec![AndConstraint([
252				vec![
253					ShiftedValueIndex::plain(public),
254					ShiftedValueIndex::srl(public, 3),
255				],
256				vec![ShiftedValueIndex::sar(hidden, 7)],
257				vec![
258					ShiftedValueIndex::rotr(hidden, 1),
259					ShiftedValueIndex::plain(hidden),
260				],
261			])],
262			imul_constraints: Vec::new(),
263			bmul_constraints: Vec::new(),
264		}
265	}
266
267	#[test]
268	fn dense_shift_encoding_covers_the_sequences_its_segment_uses() {
269		let key_collection =
270			KeyCollection::build(&shifted_constraint_system(), InoutSegment::Public);
271
272		let public_sequences = key_collection
273			.public
274			.dense_shift_enc
275			.iter()
276			.collect::<Vec<_>>();
277		assert_eq!(public_sequences, [single(Shift::IDENTITY), single(Shift::srl(3))]);
278
279		let hidden_sequences = key_collection
280			.hidden
281			.dense_shift_enc
282			.iter()
283			.collect::<Vec<_>>();
284		assert_eq!(
285			hidden_sequences,
286			[
287				single(Shift::IDENTITY),
288				single(Shift::sar(7)),
289				single(Shift::rotr(1)),
290			]
291		);
292
293		// The point of the encoding: a segment names far fewer sequences than the space holds, and
294		// the space is now the square of one slot's alphabet.
295		assert!(key_collection.hidden.dense_shift_enc.len() < SHIFT_COUNT * SHIFT_COUNT);
296	}
297
298	#[test]
299	fn dense_shift_encoding_distinguishes_sequences_sharing_an_inner_shift() {
300		// Two terms sharing an inner shift but differing outside must land on distinct indices.
301		// Keyed on the inner shift alone they would collide and accumulate into one row.
302		let hidden = ValueIndex::private(1);
303		let cs = ConstraintSystem {
304			constants: vec![Word::ZERO; 4],
305			n_inout: 0,
306			n_private: 4,
307			zero_constraints: Vec::new(),
308			and_constraints: vec![AndConstraint([
309				vec![
310					ShiftedValueIndex::new(hidden, [Shift::srl(3), Shift::sll(3)]),
311					ShiftedValueIndex::new(hidden, [Shift::srl(3), Shift::sll(5)]),
312					ShiftedValueIndex::srl(hidden, 3),
313				],
314				Vec::new(),
315				Vec::new(),
316			])],
317			imul_constraints: Vec::new(),
318			bmul_constraints: Vec::new(),
319		};
320
321		let key_collection = KeyCollection::build(&cs, InoutSegment::Public);
322		let sequences = key_collection
323			.hidden
324			.dense_shift_enc
325			.iter()
326			.collect::<Vec<_>>();
327		assert_eq!(
328			sequences,
329			[
330				single(Shift::srl(3)),
331				[Shift::srl(3), Shift::sll(3)],
332				[Shift::srl(3), Shift::sll(5)],
333			]
334		);
335
336		// The word's three keys therefore hold three distinct dense indices.
337		let mut indices = key_collection
338			.hidden
339			.word_keys(1)
340			.iter()
341			.map(|key| key.dense_shift_idx)
342			.collect::<Vec<_>>();
343		indices.sort_unstable();
344		assert_eq!(indices, [0, 1, 2]);
345	}
346
347	#[test]
348	fn keys_index_their_segments_dense_encoding() {
349		let key_collection =
350			KeyCollection::build(&shifted_constraint_system(), InoutSegment::Public);
351
352		// The shift sequences a word's keys name, as its own segment's encoding recovers them.
353		let word_sequences = |segment: &KeySegment, word: usize| {
354			let mut sequences = segment
355				.word_keys(word)
356				.iter()
357				.map(|key| {
358					segment
359						.dense_shift_enc
360						.iter()
361						.nth(key.dense_shift_idx as usize)
362						.unwrap()
363				})
364				.collect::<Vec<_>>();
365			sequences.sort();
366			sequences
367		};
368
369		// Value index 1 is the second public word; value index 5 the second hidden one.
370		assert_eq!(
371			word_sequences(&key_collection.public, 1),
372			[single(Shift::IDENTITY), single(Shift::srl(3))]
373		);
374		assert_eq!(
375			word_sequences(&key_collection.hidden, 1),
376			[
377				single(Shift::IDENTITY),
378				single(Shift::sar(7)),
379				single(Shift::rotr(1)),
380			]
381		);
382	}
383
384	#[test]
385	fn dense_shift_encoding_survives_serialization() {
386		let key_collection =
387			KeyCollection::build(&shifted_constraint_system(), InoutSegment::Public);
388
389		let mut buf = Vec::new();
390		key_collection.serialize(&mut buf).unwrap();
391		let deserialized = KeyCollection::deserialize(buf.as_slice()).unwrap();
392
393		for (segment, deserialized) in [
394			(&key_collection.public, &deserialized.public),
395			(&key_collection.hidden, &deserialized.hidden),
396		] {
397			assert_eq!(
398				segment.dense_shift_enc.iter().collect::<Vec<_>>(),
399				deserialized.dense_shift_enc.iter().collect::<Vec<_>>()
400			);
401		}
402	}
403}