openvm_deferral_circuit/poseidon2/
trace.rs

1use 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                // WARNING: Poseidon2SubCols must be the first field in DeferralPoseidon2Cols.
133                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}