openvm_recursion_circuit/stacking/eq_base/
trace.rs1use 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}