openvm_recursion_circuit/primitives/range/
trace.rs1use 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}