openvm_deferral_circuit/count/
trace.rs

1use std::{
2    borrow::BorrowMut,
3    sync::atomic::{AtomicU32, Ordering},
4};
5
6use openvm_circuit::utils::next_power_of_two_or_zero;
7use openvm_circuit_primitives::Chip;
8use openvm_cpu_backend::CpuBackend;
9use openvm_stark_backend::{
10    p3_field::PrimeCharacteristicRing, p3_matrix::dense::RowMajorMatrix, prover::AirProvingContext,
11    StarkProtocolConfig, Val,
12};
13
14use crate::count::DeferralCircuitCountCols;
15
16#[derive(Debug)]
17pub struct DeferralCircuitCountChip {
18    pub count: Vec<AtomicU32>,
19}
20
21impl DeferralCircuitCountChip {
22    pub fn new(num_deferral_circuit: usize) -> Self {
23        let count = (0..num_deferral_circuit)
24            .map(|_| AtomicU32::new(0))
25            .collect();
26        Self { count }
27    }
28
29    pub fn add_count(&self, idx: u32) {
30        let idx = idx as usize;
31        assert!(idx < self.count.len());
32        let val_atomic = &self.count[idx];
33        val_atomic.fetch_add(1, Ordering::Relaxed);
34    }
35}
36
37impl<SC: StarkProtocolConfig, RA> Chip<RA, CpuBackend<SC>> for DeferralCircuitCountChip
38where
39    Val<SC>: PrimeCharacteristicRing,
40{
41    fn constant_trace_height(&self) -> Option<usize> {
42        Some(next_power_of_two_or_zero(self.count.len()))
43    }
44
45    fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<CpuBackend<SC>> {
46        let width = DeferralCircuitCountCols::<u8>::width();
47        let height = next_power_of_two_or_zero(self.count.len());
48        let mut trace = vec![Val::<SC>::ZERO; width * height];
49
50        let mut rows = trace.chunks_exact_mut(width);
51        let mut row_idx = 0u32;
52
53        for mult in &self.count {
54            let row = rows.next().unwrap();
55            let cols: &mut DeferralCircuitCountCols<Val<SC>> = (*row).borrow_mut();
56            cols.is_valid = Val::<SC>::ONE;
57            cols.row_idx = Val::<SC>::from_u32(row_idx);
58            cols.mult = Val::<SC>::from_u32(mult.swap(0, Ordering::Relaxed));
59            row_idx += 1;
60        }
61
62        for row in rows {
63            let cols: &mut DeferralCircuitCountCols<Val<SC>> = (*row).borrow_mut();
64            cols.row_idx = Val::<SC>::from_u32(row_idx);
65            row_idx += 1;
66        }
67
68        let trace = RowMajorMatrix::new(trace, width);
69        AirProvingContext::simple_no_pis(trace)
70    }
71}