openvm_recursion_circuit/primitives/pow/
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::pow::air::PowerCheckerCols;
9
10#[derive(Debug)]
11pub struct PowerCheckerCpuTraceGenerator<const BASE: usize, const N: usize> {
12    count_pow: Vec<AtomicU32>,
13    count_range: Vec<AtomicU32>,
14}
15
16impl<const BASE: usize, const N: usize> Default for PowerCheckerCpuTraceGenerator<BASE, N> {
17    fn default() -> Self {
18        assert!(N.is_power_of_two());
19        let mut count_pow = Vec::with_capacity(N);
20        let mut count_range = Vec::with_capacity(N);
21        for _ in 0..N {
22            count_pow.push(AtomicU32::new(0));
23            count_range.push(AtomicU32::new(0));
24        }
25        Self {
26            count_pow,
27            count_range,
28        }
29    }
30}
31
32impl<const BASE: usize, const N: usize> PowerCheckerCpuTraceGenerator<BASE, N> {
33    pub fn add_pow(&self, log: usize) -> usize {
34        self.count_pow[log].fetch_add(1, Ordering::Relaxed);
35        1 << log
36    }
37
38    pub fn add_range(&self, value: usize) {
39        self.count_range[value].fetch_add(1, Ordering::Relaxed);
40    }
41
42    pub fn add_pow_count(&self, log: usize, count: u32) {
43        debug_assert!(log < self.count_pow.len());
44        if count != 0 {
45            self.count_pow[log].fetch_add(count, Ordering::Relaxed);
46        }
47    }
48
49    pub fn add_range_count(&self, value: usize, count: u32) {
50        debug_assert!(value < self.count_range.len());
51        if count != 0 {
52            self.count_range[value].fetch_add(count, Ordering::Relaxed);
53        }
54    }
55
56    pub fn take_counts(&self) -> (Vec<u32>, Vec<u32>) {
57        let pow = self
58            .count_pow
59            .iter()
60            .map(|counter| counter.swap(0, Ordering::Relaxed))
61            .collect();
62        let range = self
63            .count_range
64            .iter()
65            .map(|counter| counter.swap(0, Ordering::Relaxed))
66            .collect();
67        (pow, range)
68    }
69
70    pub fn reset(&self) {
71        for counter in &self.count_pow {
72            counter.store(0, Ordering::Relaxed);
73        }
74        for counter in &self.count_range {
75            counter.store(0, Ordering::Relaxed);
76        }
77    }
78
79    #[tracing::instrument(name = "generate_trace", level = "trace", skip_all)]
80    pub fn generate_trace_row_major(&self) -> RowMajorMatrix<F> {
81        let mut current_pow = F::ONE;
82        let trace = self
83            .count_pow
84            .iter()
85            .zip(self.count_range.iter())
86            .enumerate()
87            .flat_map(|(log, (mult_pow, mult_range))| {
88                let ret = [
89                    F::from_usize(log),
90                    current_pow,
91                    F::from_u32(mult_pow.load(Ordering::Relaxed)),
92                    F::from_u32(mult_range.load(Ordering::Relaxed)),
93                ];
94                current_pow *= F::from_usize(BASE);
95                ret
96            })
97            .collect_vec();
98        RowMajorMatrix::new(trace, PowerCheckerCols::<u8>::width())
99    }
100}