openvm_recursion_circuit/stacking/eq_bits/
trace.rs

1use 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                // (b_value, num_bits) -> (sub_eval, eval, internal_mult, external_mult)
41                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                /*
47                 * Suppose we have some b_value b[0..k], where k is total_num_bits. Then
48                 * eq_bits(u, b) is a function of eq_bits(u[0..k - 1], b[0..k - 1]), u[k],
49                 * and b[k]. This AIR uses that property to compute each eq_bits(u, b) via
50                 * a tree structure + internal interactions.
51                 */
52                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}