openvm_deferral_circuit/count/
trace.rs1use 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}