openvm_recursion_circuit/stacking/eq_base/
trace.rs

1use std::borrow::BorrowMut;
2
3use itertools::Itertools;
4use openvm_stark_backend::poly_common::{eval_eq_uni, eval_rot_kernel_prism, Squarable};
5use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
6use p3_field::{BasedVectorSpace, PrimeCharacteristicRing, TwoAdicField};
7use p3_matrix::dense::RowMajorMatrix;
8use p3_maybe_rayon::prelude::*;
9
10use crate::{
11    stacking::eq_base::air::EqBaseCols,
12    tracegen::{RowMajorChip, StandardTracegenCtx},
13};
14
15pub struct EqBaseTraceGenerator;
16
17impl RowMajorChip<F> for EqBaseTraceGenerator {
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 = EqBaseCols::<usize>::width();
32
33        let num_rows_per_proof = vk.inner.params.l_skip + 1;
34        let traces = proofs
35            .par_iter()
36            .zip(preflights.par_iter())
37            .enumerate()
38            .map(|(proof_idx, (proof, preflight))| {
39                let mut mults = vec![0usize; vk.inner.params.l_skip + 1];
40                for (sort_idx, (air_idx, vdata)) in
41                    preflight.proof_shape.sorted_trace_vdata.iter().enumerate()
42                {
43                    let need_rot = vk.inner.per_air[*air_idx].params.need_rot;
44                    if vdata.log_height <= vk.inner.params.l_skip {
45                        let neg_n = vk.inner.params.l_skip - vdata.log_height;
46                        mults[neg_n] += proof.batch_constraint_proof.column_openings[sort_idx]
47                            .iter()
48                            .flatten()
49                            .count()
50                            / if need_rot { 2 } else { 1 };
51                    }
52                }
53
54                let proof_idx_value = F::from_usize(proof_idx);
55
56                let mut trace = vec![F::ZERO; num_rows_per_proof * width];
57
58                let omega = F::two_adic_generator(vk.inner.params.l_skip);
59                let mut u = preflight.stacking.sumcheck_rnd[0];
60                let mut r = preflight.batch_constraint.sumcheck_rnd[0];
61                let mut r_omega = r * omega;
62
63                let mut prod_u_r = u * (u + r);
64                let mut prod_u_r_omega = u * (u + r_omega);
65                let mut prod_u_1 = u + F::ONE;
66                let mut prod_r_omega_1 = r_omega + F::ONE;
67
68                let mut in_prod = EF::ONE;
69
70                let u_pows = u
71                    .exp_powers_of_2()
72                    .take(vk.inner.params.l_skip + 1)
73                    .collect_vec();
74
75                for (row_idx, chunk) in trace.chunks_mut(width).take(num_rows_per_proof).enumerate()
76                {
77                    let cols: &mut EqBaseCols<F> = chunk.borrow_mut();
78                    let is_last = row_idx + 1 == num_rows_per_proof;
79
80                    cols.proof_idx = proof_idx_value;
81                    cols.is_valid = F::ONE;
82                    cols.is_first = F::from_bool(row_idx == 0);
83                    cols.is_last = F::from_bool(is_last);
84
85                    cols.row_idx = F::from_usize(row_idx);
86
87                    cols.u_pow.copy_from_slice(u.as_basis_coefficients_slice());
88                    cols.r_pow.copy_from_slice(r.as_basis_coefficients_slice());
89                    cols.r_omega_pow
90                        .copy_from_slice(r_omega.as_basis_coefficients_slice());
91
92                    cols.prod_u_r
93                        .copy_from_slice(prod_u_r.as_basis_coefficients_slice());
94                    cols.prod_u_r_omega
95                        .copy_from_slice(prod_u_r_omega.as_basis_coefficients_slice());
96                    cols.prod_u_1
97                        .copy_from_slice(prod_u_1.as_basis_coefficients_slice());
98                    cols.prod_r_omega_1
99                        .copy_from_slice(prod_r_omega_1.as_basis_coefficients_slice());
100
101                    if is_last {
102                        cols.mult = F::from_usize(mults[0]);
103                    }
104
105                    let l_skip = vk.inner.params.l_skip - row_idx;
106                    let u_pow_rev = u_pows[l_skip];
107
108                    if row_idx != 0 {
109                        in_prod *= u_pow_rev + F::ONE;
110                        cols.eq_neg.copy_from_slice(
111                            (eval_eq_uni(l_skip, preflight.stacking.sumcheck_rnd[0], r)
112                                * F::from_usize(1 << l_skip))
113                            .as_basis_coefficients_slice(),
114                        );
115                        cols.k_rot_neg.copy_from_slice(
116                            (eval_rot_kernel_prism(
117                                l_skip,
118                                &[preflight.stacking.sumcheck_rnd[0]],
119                                &[r],
120                            ) * F::from_usize(1 << l_skip))
121                            .as_basis_coefficients_slice(),
122                        );
123                        cols.mult_neg = F::from_usize(mults[row_idx]);
124                    }
125
126                    cols.u_pow_rev
127                        .copy_from_slice(u_pow_rev.as_basis_coefficients_slice());
128                    cols.in_prod
129                        .copy_from_slice(in_prod.as_basis_coefficients_slice());
130
131                    u *= u;
132                    r *= r;
133                    r_omega *= r_omega;
134
135                    prod_u_r *= u + r;
136                    prod_u_r_omega *= u + r_omega;
137                    prod_u_1 *= u + F::ONE;
138                    prod_r_omega_1 *= r_omega + F::ONE;
139                }
140
141                (trace, num_rows_per_proof)
142            })
143            .collect::<Vec<_>>();
144
145        let num_valid_rows = traces.iter().map(|(_trace, num_rows)| *num_rows).sum();
146        let height = if let Some(height) = required_height {
147            if height < num_valid_rows {
148                return None;
149            }
150            height
151        } else {
152            num_valid_rows.next_power_of_two()
153        };
154
155        let mut combined_trace = Vec::with_capacity(height * width);
156        for (trace, _num_rows) in traces {
157            combined_trace.extend(trace);
158        }
159
160        let padding_proof_idx = F::from_usize(proofs.len());
161        combined_trace.resize(height * width, F::ZERO);
162        let mut chunks = combined_trace[num_valid_rows * width..]
163            .chunks_mut(width)
164            .peekable();
165
166        while let Some(chunk) = chunks.next() {
167            let cols: &mut EqBaseCols<F> = chunk.borrow_mut();
168            cols.proof_idx = padding_proof_idx;
169            if chunks.peek().is_none() {
170                cols.is_last = F::ONE;
171            }
172        }
173
174        Some(RowMajorMatrix::new(combined_trace, width))
175    }
176}