openvm_recursion_circuit/primitives/range/
trace.rs

1use std::sync::atomic::{AtomicU32, Ordering};
2
3use itertools::Itertools;
4use openvm_stark_sdk::config::baby_bear_poseidon2::F;
5use p3_field::PrimeCharacteristicRing;
6use p3_matrix::dense::RowMajorMatrix;
7
8use crate::primitives::range::air::RangeCheckerCols;
9
10#[derive(Debug)]
11pub struct RangeCheckerCpuTraceGenerator<const NUM_BITS: usize> {
12    count: Vec<AtomicU32>,
13}
14
15impl<const NUM_BITS: usize> Default for RangeCheckerCpuTraceGenerator<NUM_BITS> {
16    fn default() -> Self {
17        let mut count = Vec::with_capacity(1 << NUM_BITS);
18        for _ in 0..(1 << NUM_BITS) {
19            count.push(AtomicU32::new(0));
20        }
21        Self { count }
22    }
23}
24
25impl<const NUM_BITS: usize> RangeCheckerCpuTraceGenerator<NUM_BITS> {
26    pub fn add_count(&self, value: usize) {
27        self.add_count_mult(value, 1);
28    }
29
30    pub fn add_count_mult(&self, value: usize, mult: u32) {
31        self.count[value].fetch_add(mult, Ordering::Relaxed);
32    }
33
34    #[tracing::instrument(name = "generate_trace", level = "trace", skip_all)]
35    pub fn generate_trace_row_major(&self) -> RowMajorMatrix<F> {
36        let trace = self
37            .count
38            .iter()
39            .enumerate()
40            .flat_map(|(value, mult)| {
41                [
42                    F::from_usize(value),
43                    F::from_u32(mult.load(Ordering::Relaxed)),
44                ]
45            })
46            .collect_vec();
47        RowMajorMatrix::new(trace, RangeCheckerCols::<u8>::width())
48    }
49}