openvm_deferral_circuit/poseidon2/
trace.rs1use std::{
2 array::from_fn,
3 borrow::BorrowMut,
4 sync::atomic::{AtomicBool, AtomicU32},
5};
6
7use dashmap::DashMap;
8use openvm_circuit::arch::VmField;
9use openvm_circuit_primitives::{utils::next_power_of_two_or_zero, Chip};
10use openvm_cpu_backend::CpuBackend;
11use openvm_poseidon2_air::{Poseidon2Config, Poseidon2SubChip, POSEIDON2_WIDTH};
12use openvm_stark_backend::{
13 p3_air::BaseAir, p3_field::PrimeCharacteristicRing, p3_matrix::dense::RowMajorMatrix,
14 p3_maybe_rayon::prelude::*, prover::AirProvingContext, StarkProtocolConfig, Val,
15};
16use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
17use rustc_hash::FxBuildHasher;
18
19use super::{DeferralPoseidon2Cols, SBOX_REGISTERS};
20use crate::chunks_to_state;
21
22#[derive(Debug)]
23pub struct DeferralPoseidon2Chip<F: VmField> {
24 pub subchip: Poseidon2SubChip<F, SBOX_REGISTERS>,
25 pub records: DashMap<[F; POSEIDON2_WIDTH], (AtomicU32, AtomicU32), FxBuildHasher>,
26 pub nonempty: AtomicBool,
27}
28
29impl<F: VmField> DeferralPoseidon2Chip<F> {
30 pub fn new(poseidon2_config: Poseidon2Config<F>) -> Self {
31 let subchip = Poseidon2SubChip::new(poseidon2_config.constants);
32 Self {
33 subchip,
34 records: DashMap::default(),
35 nonempty: AtomicBool::new(false),
36 }
37 }
38
39 pub fn perm(
40 &self,
41 lhs: &[F; DIGEST_SIZE],
42 rhs: &[F; DIGEST_SIZE],
43 is_compress: bool,
44 ) -> [F; DIGEST_SIZE] {
45 let output = self.perm_state(lhs, rhs);
46 self.select_output_chunk(output, is_compress)
47 }
48
49 pub fn perm_state(
50 &self,
51 lhs: &[F; DIGEST_SIZE],
52 rhs: &[F; DIGEST_SIZE],
53 ) -> [F; POSEIDON2_WIDTH] {
54 let input = chunks_to_state(lhs, rhs);
55 self.subchip.permute(input)
56 }
57
58 pub fn perm_and_record(
59 &self,
60 lhs: &[F; DIGEST_SIZE],
61 rhs: &[F; DIGEST_SIZE],
62 is_compress: bool,
63 ) -> [F; DIGEST_SIZE] {
64 let input = chunks_to_state(lhs, rhs);
65 let output = self.subchip.permute(input);
66 let ret = self.select_output_chunk(output, is_compress);
67 let count = self
68 .records
69 .entry(input)
70 .or_insert((AtomicU32::new(0), AtomicU32::new(0)));
71 let mult = if is_compress { &count.0 } else { &count.1 };
72 mult.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
73 self.nonempty
74 .store(true, std::sync::atomic::Ordering::Relaxed);
75 ret
76 }
77
78 fn select_output_chunk(
79 &self,
80 output: [F; POSEIDON2_WIDTH],
81 is_compress: bool,
82 ) -> [F; DIGEST_SIZE] {
83 let offset = if is_compress { 0 } else { DIGEST_SIZE };
84 from_fn(|i| output[i + offset])
85 }
86}
87
88impl<RA, SC: StarkProtocolConfig> Chip<RA, CpuBackend<SC>> for DeferralPoseidon2Chip<Val<SC>>
89where
90 Val<SC>: VmField,
91{
92 fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<CpuBackend<SC>> {
93 let width = DeferralPoseidon2Cols::<Val<SC>>::width();
94 if !self.nonempty.load(std::sync::atomic::Ordering::Relaxed) {
95 let trace = RowMajorMatrix::new(vec![], width);
96 return AirProvingContext::simple_no_pis(trace);
97 }
98 let height = next_power_of_two_or_zero(self.records.len());
99
100 let mut inputs = Vec::with_capacity(height);
101 let mut multiplicities = Vec::with_capacity(height);
102 #[cfg(feature = "parallel")]
103 let records_iter = self.records.par_iter();
104 #[cfg(not(feature = "parallel"))]
105 let records_iter = self.records.iter();
106 let (actual_inputs, actual_multiplicities): (Vec<_>, Vec<_>) = records_iter
107 .map(|record| {
108 let (input, (compress_mult, capacity_mult)) = record.pair();
109 (
110 *input,
111 (
112 compress_mult.load(std::sync::atomic::Ordering::Relaxed),
113 capacity_mult.load(std::sync::atomic::Ordering::Relaxed),
114 ),
115 )
116 })
117 .unzip();
118 inputs.extend(actual_inputs);
119 multiplicities.extend(actual_multiplicities);
120 inputs.resize(height, [Val::<SC>::ZERO; POSEIDON2_WIDTH]);
121 multiplicities.resize(height, (0, 0));
122
123 let inner_trace = self.subchip.generate_trace(inputs);
124 let inner_width = self.subchip.air.width();
125
126 let mut values = Val::<SC>::zero_vec(height * width);
127 values
128 .par_chunks_mut(width)
129 .zip(inner_trace.values.par_chunks(inner_width))
130 .zip(multiplicities)
131 .for_each(|((row, inner_row), (compress_mult, capacity_mult))| {
132 row[..inner_width].copy_from_slice(inner_row);
134 let cols: &mut DeferralPoseidon2Cols<Val<SC>> = row.borrow_mut();
135 cols.compress_mult = Val::<SC>::from_u32(compress_mult);
136 cols.capacity_mult = Val::<SC>::from_u32(capacity_mult);
137 });
138 self.records.clear();
139 self.nonempty
140 .store(false, std::sync::atomic::Ordering::Relaxed);
141
142 let trace = RowMajorMatrix::new(values, width);
143 AirProvingContext::simple_no_pis(trace)
144 }
145}