openvm_recursion_circuit/primitives/pow/
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::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}