binius_prover/protocols/shift/key_collection/
key_segment.rs1use 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#[derive(Debug, Clone)]
29pub struct KeySegment {
30 pub keys: Vec<Key>,
32 pub key_ranges: Vec<Range<u32>>,
35 pub constraint_indices: Vec<ConstraintIndex>,
37 pub dense_shift_enc: DenseShiftEncoding,
39}
40
41impl KeySegment {
42 pub const fn n_words(&self) -> usize {
44 self.key_ranges.len()
45 }
46
47 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 #[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 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 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 let low_bits_mask = (1u8 << P::WIDTH) - 1;
88
89 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 for key in keys {
105 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 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 let packed_mask =
142 packed_masks_map.get_unchecked(packed_mask_index);
143
144 *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 .reduce_with(|mut acc, local| {
166 izip!(acc.iter_mut(), local.iter()).for_each(|(acc, local)| {
167 *acc += *local;
168 });
169 acc
170 })
171 .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 (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 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 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 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 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 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 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 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 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 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}