binius_prover/protocols/shift/key_collection/
collection.rs1use 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#[derive(Debug, Clone)]
31pub struct KeyCollection {
32 pub public: KeySegment,
34 pub hidden: KeySegment,
36}
37
38impl KeyCollection {
39 pub fn build(cs: &ConstraintSystem, inout: InoutSegment) -> Self {
50 builder::build_key_collection(cs, inout)
51 }
52
53 pub const fn log_witness_words(&self) -> usize {
62 log2_ceil_usize(self.hidden.n_words())
63 }
64
65 #[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 let outer_weights = OuterSlotWeights::<F>::new(outer);
101
102 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 let word_scalar = |segment: &KeySegment, scalars: &[F], index: usize| {
127 let wide = segment
128 .word_keys(index)
129 .iter()
130 .map(|key| {
131 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 let build_segment = |segment: &KeySegment, log_len: usize| {
148 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 let n_full = n_words / P::WIDTH;
156 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 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 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 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 fn single(shift: Shift) -> [Shift; 2] {
233 [shift, Shift::IDENTITY]
234 }
235
236 fn shifted_constraint_system() -> ConstraintSystem {
242 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 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 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 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 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 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}