Skip to main content

binius_prover/protocols/shift/key_collection/
key_segment.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::ops::Range;
5
6use binius_core::word::Word;
7use binius_field::{Field, PackedField};
8use binius_utils::{
9	rayon::{
10		prelude::*,
11		task_size::{IndexedParallelIteratorExt, WorkPerItem},
12	},
13	serialization::{DeserializeBytes, SerializationError, SerializeBytes},
14};
15use bytemuck::zeroed_vec;
16use bytes::{Buf, BufMut};
17use itertools::izip;
18use tracing::instrument;
19
20use super::{
21	super::{claims::PreparedOperandClaims, phase_1::row_len},
22	dense_shift_encoding::DenseShiftEncoding,
23	key::{ConstraintIndex, Key},
24};
25
26/// One value-vector segment's keys, public or hidden.
27/// Indexed so each word's constraints can be found without a scan.
28#[derive(Debug, Clone)]
29pub struct KeySegment {
30	/// Every key of the segment, flattened into one vector.
31	pub keys: Vec<Key>,
32	/// One range per word, at that word's segment-relative index.
33	/// The range names that word's keys inside the flattened keys vector.
34	pub key_ranges: Vec<Range<u32>>,
35	/// The constraint indices the keys reference, flattened into one vector.
36	pub constraint_indices: Vec<ConstraintIndex>,
37	/// The shift sequences the segment's keys name.
38	pub dense_shift_enc: DenseShiftEncoding,
39}
40
41impl KeySegment {
42	/// The number of words the segment covers.
43	pub const fn n_words(&self) -> usize {
44		self.key_ranges.len()
45	}
46
47	/// The keys for the word at the given segment-relative index.
48	pub fn word_keys(&self, index: usize) -> &[Key] {
49		let Range { start, end } = self.key_ranges[index];
50		&self.keys[start as usize..end as usize]
51	}
52
53	/// Accumulates this segment's rows of the witness-and-batching multilinear.
54	///
55	/// Every witness word can carry several shift keys, one per operand position it feeds.
56	/// For each key, this folds its constraint-index tensor and batching powers into one
57	/// accumulator value, then scatters that value across the bit positions the word has set.
58	///
59	/// # Returns
60	///
61	/// One row of weights per shift this segment uses, one weight per bit position of a word,
62	/// at the position this segment's own dense encoding assigns to that shift.
63	#[instrument(skip_all, name = "build_g")]
64	pub fn build_g<F: Field, P: PackedField<Scalar = F>>(
65		&self,
66		words: &[Word],
67		prepared: &PreparedOperandClaims<F>,
68	) -> Box<[P]> {
69		// One row of `Word::BITS` scalars per shift the segment uses.
70		let row_len = row_len::<P>();
71		let acc_size = self.dense_shift_enc.len() * row_len;
72
73		const {
74			assert!(
75				P::WIDTH <= 8,
76				"the optimizations below work only when the width of `P` is less than 8 (which is true for all packed 128b fields we use for now)"
77			);
78		}
79
80		// Precompute once: a map from an 8-bit mask (one bit per candidate lane) to the packed
81		// lane-selection mask it names.
82		// Every accumulator below reuses this table instead of rebuilding a mask per word.
83		let packed_masks_map = (0..1 << P::WIDTH)
84			.map(|i| P::make_mask((0..P::WIDTH).map(|bit_index| (i >> bit_index) & 1 == 1)))
85			.collect::<Vec<_>>();
86		// A mask for the low bits that actually address a lane.
87		let low_bits_mask = (1u8 << P::WIDTH) - 1;
88
89		// Each word carries the keys its segment-relative range names.
90		//
91		// The fold runs in parallel.
92		// Each task owns a chunk of words and its own accumulator.
93		// The accumulators merge at the end.
94		words
95			.par_iter()
96			.zip(self.key_ranges.par_iter())
97			.with_min_task(WorkPerItem::FieldMuls)
98			.fold(
99				|| zeroed_vec::<P>(acc_size).into_boxed_slice(),
100				|mut multilinears, (word, Range { start, end })| {
101					let keys = &self.keys[*start as usize..*end as usize];
102
103					// A word can carry several keys, one per shift sequence it is read under.
104					for key in keys {
105						// Fold this key's accumulator value: the constraint-index tensor
106						// against the operand weight of each column the key names.
107						let acc = key.accumulate(
108							&self.constraint_indices,
109							&prepared.r_x_tensor,
110							&prepared.operand_weights,
111						);
112						let acc_packed = P::broadcast(acc);
113
114						// Scatter the accumulator value into every set bit of the word.
115						//
116						// This walks only the word's nonzero bytes, instead of testing all
117						// `Word::BITS` positions in turn.
118						// Within each byte, it selects every set bit's lane in one packed
119						// operation, using the precomputed mask above.
120						debug_assert!(
121							(key.dense_shift_idx as usize) < self.dense_shift_enc.len(),
122							"a key indexes the dense shift encoding of its own segment"
123						);
124						let start = key.dense_shift_idx as usize * row_len;
125						let values = &mut multilinears[start..start + row_len];
126						let values_per_byte = Word::BYTES >> P::LOG_WIDTH;
127						let mut remaining_word = word.0;
128						let mut byte_index = 0;
129						while remaining_word != 0 {
130							let byte = remaining_word as u8;
131							let byte_values =
132								&mut values[byte_index * values_per_byte..][..values_per_byte];
133							for value_index in 0..(8 >> P::LOG_WIDTH) {
134								unsafe {
135									let packed_mask_index = ((byte >> (value_index * P::WIDTH))
136										& low_bits_mask) as usize;
137
138									// Safety:
139									// - `packed_masks_map` holds one entry per possible
140									//   `P::WIDTH`-bit value, so this index always fits.
141									let packed_mask =
142										packed_masks_map.get_unchecked(packed_mask_index);
143
144									// Safety:
145									// - `values` spans `(8 >> P::LOG_WIDTH)` elements after the
146									//   chunking above.
147									// - `value_index` stays within that range by construction.
148									*byte_values.get_unchecked_mut(value_index) +=
149										acc_packed.select(packed_mask);
150								}
151							}
152							remaining_word >>= 8;
153							byte_index += 1;
154						}
155					}
156
157					multilinears
158				},
159			)
160			// Merge every task's accumulator pairwise.
161			//
162			// A merge seeded with an existing partial never touches a zeroed buffer.
163			// An identity element instead would allocate and zero one accumulator per merge.
164			// Then it would add all of it for nothing.
165			.reduce_with(|mut acc, local| {
166				izip!(acc.iter_mut(), local.iter()).for_each(|(acc, local)| {
167					*acc += *local;
168				});
169				acc
170			})
171			// An empty word list produces no partial accumulator at all.
172			// So its rows are zero.
173			.unwrap_or_else(|| zeroed_vec::<P>(acc_size).into_boxed_slice())
174	}
175}
176
177impl SerializeBytes for KeySegment {
178	fn serialize(&self, mut write_buf: impl BufMut) -> Result<(), SerializationError> {
179		self.keys.serialize(&mut write_buf)?;
180
181		// Serialize key_ranges as pairs of start/end
182		(self.key_ranges.len() as u32).serialize(&mut write_buf)?;
183		for range in &self.key_ranges {
184			range.start.serialize(&mut write_buf)?;
185			range.end.serialize(&mut write_buf)?;
186		}
187
188		self.constraint_indices.serialize(&mut write_buf)?;
189		self.dense_shift_enc.serialize(write_buf)
190	}
191}
192
193impl DeserializeBytes for KeySegment {
194	fn deserialize(mut read_buf: impl Buf) -> Result<Self, SerializationError> {
195		let keys = Vec::<Key>::deserialize(&mut read_buf)?;
196
197		// Deserialize key_ranges
198		let len = u32::deserialize(&mut read_buf)? as usize;
199		let mut key_ranges = Vec::with_capacity(len);
200		for _ in 0..len {
201			let start = u32::deserialize(&mut read_buf)?;
202			let end = u32::deserialize(&mut read_buf)?;
203			key_ranges.push(start..end);
204		}
205
206		let constraint_indices = Vec::<ConstraintIndex>::deserialize(&mut read_buf)?;
207		let dense_shift_enc = DenseShiftEncoding::deserialize(&mut read_buf)?;
208
209		// A key indexes its own segment's shift encoding and constraint list.
210		// Rejecting a bad index here beats panicking mid-proof in `build_g` or `word_scalar`.
211		let keys_index_siblings = keys.iter().all(|key| {
212			(key.dense_shift_idx as usize) < dense_shift_enc.len()
213				&& key.range.end as usize <= constraint_indices.len()
214		});
215		if !keys_index_siblings {
216			return Err(SerializationError::InvalidConstruction {
217				name: "KeySegment::keys",
218			});
219		}
220		// A word's range names its keys inside the flattened keys vector, which `word_keys` slices.
221		let key_ranges_index_keys = key_ranges
222			.iter()
223			.all(|range| range.start <= range.end && range.end as usize <= keys.len());
224		if !key_ranges_index_keys {
225			return Err(SerializationError::InvalidConstruction {
226				name: "KeySegment::key_ranges",
227			});
228		}
229
230		Ok(KeySegment {
231			keys,
232			key_ranges,
233			constraint_indices,
234			dense_shift_enc,
235		})
236	}
237}
238
239#[cfg(test)]
240mod tests {
241	use binius_core::constraint_system::Shift;
242
243	use super::*;
244
245	// Serializes a segment built raw, bypassing `build`, so malformed indices reach the
246	// deserializer.
247	fn deserialize_raw(segment: &KeySegment) -> Result<KeySegment, SerializationError> {
248		let mut buf = Vec::new();
249		segment.serialize(&mut buf).unwrap();
250		KeySegment::deserialize(buf.as_slice())
251	}
252
253	#[test]
254	fn key_segment_round_trips_a_well_formed_segment() {
255		// Pins the checks against the shape `build` produces: word ranges tile the keys vector, and
256		// every index addresses a sibling long enough to hold it.
257		let segment = KeySegment {
258			keys: vec![
259				Key {
260					dense_shift_idx: 0,
261					range: Range { start: 0, end: 2 },
262				},
263				Key {
264					dense_shift_idx: 1,
265					range: Range { start: 2, end: 3 },
266				},
267			],
268			key_ranges: vec![Range { start: 0, end: 1 }, Range { start: 1, end: 2 }],
269			constraint_indices: (0..3)
270				.map(|constraint_index| ConstraintIndex {
271					operand_index: 0,
272					constraint_index,
273				})
274				.collect(),
275			dense_shift_enc: DenseShiftEncoding::new([
276				[Shift::srl(1), Shift::IDENTITY],
277				[Shift::srl(2), Shift::IDENTITY],
278			]),
279		};
280
281		let segment = deserialize_raw(&segment).expect("a well-formed segment deserializes");
282
283		assert_eq!(segment.n_words(), 2);
284		assert_eq!(segment.keys.len(), 2);
285		assert_eq!(segment.constraint_indices.len(), 3);
286		assert_eq!(segment.dense_shift_enc.len(), 2);
287		assert_eq!(segment.word_keys(0).len(), 1);
288		assert_eq!(segment.word_keys(1).len(), 1);
289	}
290
291	#[test]
292	fn key_segment_rejects_a_key_indexing_past_the_shift_encoding() {
293		// `build_g` scales this index by the row length to reach into the multilinears buffer.
294		// The encoding holds two sequences, so index 2 is one past its end.
295		let segment = KeySegment {
296			keys: vec![Key {
297				dense_shift_idx: 2,
298				range: Range { start: 0, end: 1 },
299			}],
300			key_ranges: vec![Range { start: 0, end: 1 }],
301			constraint_indices: vec![ConstraintIndex {
302				operand_index: 0,
303				constraint_index: 0,
304			}],
305			dense_shift_enc: DenseShiftEncoding::new([
306				[Shift::srl(1), Shift::IDENTITY],
307				[Shift::srl(2), Shift::IDENTITY],
308			]),
309		};
310
311		match deserialize_raw(&segment).unwrap_err() {
312			SerializationError::InvalidConstruction { name } => {
313				assert_eq!(name, "KeySegment::keys");
314			}
315			other => panic!("Expected InvalidConstruction, got: {other:?}"),
316		}
317	}
318
319	#[test]
320	fn key_segment_rejects_a_key_range_past_the_constraint_indices() {
321		// `accumulate_wide` slices the segment's flattened constraint list with this range, which
322		// holds one entry against the key's claim of four.
323		let segment = KeySegment {
324			keys: vec![Key {
325				dense_shift_idx: 0,
326				range: Range { start: 0, end: 4 },
327			}],
328			key_ranges: vec![Range { start: 0, end: 1 }],
329			constraint_indices: vec![ConstraintIndex {
330				operand_index: 0,
331				constraint_index: 0,
332			}],
333			dense_shift_enc: DenseShiftEncoding::new([[Shift::srl(1), Shift::IDENTITY]]),
334		};
335
336		match deserialize_raw(&segment).unwrap_err() {
337			SerializationError::InvalidConstruction { name } => {
338				assert_eq!(name, "KeySegment::keys");
339			}
340			other => panic!("Expected InvalidConstruction, got: {other:?}"),
341		}
342	}
343
344	#[test]
345	fn key_segment_rejects_a_word_range_past_the_keys() {
346		// `word_keys` slices the flattened keys vector, which holds one key against a range of two.
347		let segment = KeySegment {
348			keys: vec![Key {
349				dense_shift_idx: 0,
350				range: Range { start: 0, end: 1 },
351			}],
352			key_ranges: vec![Range { start: 0, end: 2 }],
353			constraint_indices: vec![ConstraintIndex {
354				operand_index: 0,
355				constraint_index: 0,
356			}],
357			dense_shift_enc: DenseShiftEncoding::new([[Shift::srl(1), Shift::IDENTITY]]),
358		};
359
360		match deserialize_raw(&segment).unwrap_err() {
361			SerializationError::InvalidConstruction { name } => {
362				assert_eq!(name, "KeySegment::key_ranges");
363			}
364			other => panic!("Expected InvalidConstruction, got: {other:?}"),
365		}
366	}
367
368	#[test]
369	fn key_segment_rejects_a_reversed_word_range() {
370		// Slicing `start..end` panics when start runs past end, in bounds or not.
371		let segment = KeySegment {
372			keys: vec![Key {
373				dense_shift_idx: 0,
374				range: Range { start: 0, end: 1 },
375			}],
376			key_ranges: vec![Range { start: 1, end: 0 }],
377			constraint_indices: vec![ConstraintIndex {
378				operand_index: 0,
379				constraint_index: 0,
380			}],
381			dense_shift_enc: DenseShiftEncoding::new([[Shift::srl(1), Shift::IDENTITY]]),
382		};
383
384		match deserialize_raw(&segment).unwrap_err() {
385			SerializationError::InvalidConstruction { name } => {
386				assert_eq!(name, "KeySegment::key_ranges");
387			}
388			other => panic!("Expected InvalidConstruction, got: {other:?}"),
389		}
390	}
391}