openvm_recursion_circuit/stacking/eq_bits/
trace.rs1use std::{borrow::BorrowMut, collections::HashMap};
2
3#[cfg(all(test, feature = "cuda"))]
4use itertools::Itertools;
5use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
6use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
7use p3_matrix::dense::RowMajorMatrix;
8use p3_maybe_rayon::prelude::*;
9
10use crate::{
11 stacking::{eq_bits::air::EqBitsCols, utils::get_stacked_slice_data},
12 tracegen::{RowMajorChip, StandardTracegenCtx},
13};
14
15pub struct EqBitsTraceGenerator;
16
17impl RowMajorChip<F> for EqBitsTraceGenerator {
18 type Ctx<'a> = StandardTracegenCtx<'a>;
19
20 #[tracing::instrument(level = "trace", skip_all)]
21 fn generate_trace(
22 &self,
23 ctx: &Self::Ctx<'_>,
24 required_height: Option<usize>,
25 ) -> Option<RowMajorMatrix<F>> {
26 let vk = ctx.vk;
27 let proofs = ctx.proofs;
28 let preflights = ctx.preflights;
29 debug_assert_eq!(proofs.len(), preflights.len());
30
31 let width = EqBitsCols::<usize>::width();
32
33 let traces = preflights
34 .par_iter()
35 .enumerate()
36 .map(|(proof_idx, preflight)| {
37 let stacked_slices =
38 get_stacked_slice_data(vk, &preflight.proof_shape.sorted_trace_vdata);
39
40 let mut b_value_map = HashMap::<(usize, usize), (EF, EF, usize, usize)>::new();
42 let mut base_internal_mult = 0usize;
43 let mut base_external_mult = 0usize;
44 let u = &preflight.stacking.sumcheck_rnd[1..];
45
46 for slice in stacked_slices {
53 let n_lift = slice.n.max(0) as usize;
54 let b_value = slice.row_idx >> (n_lift + vk.inner.params.l_skip);
55 let total_num_bits = vk.inner.params.n_stack - n_lift;
56
57 if total_num_bits == 0 {
58 base_external_mult += 1;
59 continue;
60 }
61
62 let (mut latest_eval, latest_num_bits) = {
63 let mut ret = (EF::ONE, 0);
64 for num_bits in (1..=total_num_bits).rev() {
65 let shifted_b_value = b_value >> (total_num_bits - num_bits);
66 if let Some((_, eval, internal_mult, external_mult)) =
67 b_value_map.get_mut(&(shifted_b_value, num_bits))
68 {
69 if num_bits < total_num_bits {
70 let child_b_value = b_value >> (total_num_bits - num_bits - 1);
71 *internal_mult += 1 + (child_b_value & 1);
72 } else {
73 *external_mult += 1;
74 }
75 ret = (*eval, num_bits);
76 break;
77 }
78 }
79 ret
80 };
81
82 if latest_num_bits == total_num_bits {
83 continue;
84 } else if latest_num_bits == 0 {
85 let b_value_msb = b_value >> (total_num_bits - 1);
86 base_internal_mult += 1 + b_value_msb;
87 }
88
89 for num_bits in latest_num_bits + 1..=total_num_bits {
90 let shifted_b_value = b_value >> (total_num_bits - num_bits);
91 let b_lsb = EF::from_usize(shifted_b_value & 1);
92 let u_val = u[vk.inner.params.n_stack - num_bits];
93 let next_eval =
94 latest_eval * (EF::ONE + EF::TWO * b_lsb * u_val - b_lsb - u_val);
95 let is_last = num_bits == total_num_bits;
96 b_value_map.insert(
97 (shifted_b_value, num_bits),
98 (latest_eval, next_eval, !is_last as usize, is_last as usize),
99 );
100 latest_eval = next_eval;
101 }
102 }
103
104 let num_rows = b_value_map.len() + 1;
105 let proof_idx_value = F::from_usize(proof_idx);
106
107 let mut trace = vec![F::ZERO; num_rows * width];
108
109 {
110 let first_cols: &mut EqBitsCols<F> = trace[..width].borrow_mut();
111 first_cols.proof_idx = proof_idx_value;
112 first_cols.is_valid = F::ONE;
113 first_cols.is_first = F::ONE;
114
115 first_cols.sub_eval[0] = F::ONE;
116
117 first_cols.internal_child_flag = F::from_usize(base_internal_mult);
118 first_cols.external_mult = F::from_usize(base_external_mult);
119 }
120
121 #[cfg(all(test, feature = "cuda"))]
122 let b_value_iter = b_value_map.iter().sorted();
123 #[cfg(any(not(test), not(feature = "cuda")))]
124 let b_value_iter = b_value_map.iter();
125
126 for ((&(b_value, num_bits), &(sub_eval, _, internal_mult, external_mult)), chunk) in
127 b_value_iter.zip(trace.chunks_mut(width).skip(1).take(b_value_map.len()))
128 {
129 let cols: &mut EqBitsCols<F> = chunk.borrow_mut();
130 cols.proof_idx = proof_idx_value;
131 cols.is_valid = F::ONE;
132
133 cols.internal_child_flag = F::from_usize(internal_mult);
134 cols.external_mult = F::from_usize(external_mult);
135
136 cols.sub_b_value = F::from_usize(b_value >> 1);
137 cols.num_bits = F::from_usize(num_bits);
138
139 cols.b_lsb = F::from_usize(b_value & 1);
140 cols.u_val.copy_from_slice(
141 u[vk.inner.params.n_stack - num_bits].as_basis_coefficients_slice(),
142 );
143 cols.sub_eval
144 .copy_from_slice(sub_eval.as_basis_coefficients_slice());
145 }
146
147 (trace, num_rows)
148 })
149 .collect::<Vec<_>>();
150
151 let num_valid_rows = traces.iter().map(|(_trace, num_rows)| *num_rows).sum();
152 let height = if let Some(height) = required_height {
153 if height < num_valid_rows {
154 return None;
155 }
156 height
157 } else {
158 num_valid_rows.next_power_of_two()
159 };
160
161 let mut combined_trace = Vec::with_capacity(height * width);
162 for (trace, _num_rows) in traces {
163 combined_trace.extend(trace);
164 }
165
166 let padding_proof_idx = F::from_usize(proofs.len());
167 combined_trace.resize(height * width, F::ZERO);
168 for chunk in combined_trace[num_valid_rows * width..].chunks_mut(width) {
169 let cols: &mut EqBitsCols<F> = chunk.borrow_mut();
170 cols.proof_idx = padding_proof_idx;
171 }
172
173 Some(RowMajorMatrix::new(combined_trace, width))
174 }
175}