openvm_recursion_circuit/stacking/sumcheck/
trace.rs1use std::{borrow::BorrowMut, collections::HashSet};
2
3use itertools::izip;
4use openvm_stark_backend::poly_common::{
5 eval_eq_uni, eval_eq_uni_at_one, interpolate_quadratic_at_012,
6};
7use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, EF, F};
8use p3_field::{BasedVectorSpace, PrimeCharacteristicRing, TwoAdicField};
9use p3_matrix::dense::RowMajorMatrix;
10use p3_maybe_rayon::prelude::*;
11
12use crate::{
13 stacking::{sumcheck::air::SumcheckRoundsCols, utils::get_stacked_slice_data},
14 tracegen::{RowMajorChip, StandardTracegenCtx},
15};
16
17pub struct SumcheckRoundsTraceGenerator;
18
19impl RowMajorChip<F> for SumcheckRoundsTraceGenerator {
20 type Ctx<'a> = StandardTracegenCtx<'a>;
21
22 #[tracing::instrument(level = "trace", skip_all)]
23 fn generate_trace(
24 &self,
25 ctx: &Self::Ctx<'_>,
26 required_height: Option<usize>,
27 ) -> Option<RowMajorMatrix<F>> {
28 let vk = ctx.vk;
29 let proofs = ctx.proofs;
30 let preflights = ctx.preflights;
31 debug_assert_eq!(proofs.len(), preflights.len());
32
33 let width = SumcheckRoundsCols::<usize>::width();
34
35 let traces = proofs
36 .par_iter()
37 .zip(preflights.par_iter())
38 .enumerate()
39 .map(|(proof_idx, (proof, preflight))| {
40 let sumcheck_rounds = &proof.stacking_proof.sumcheck_round_polys;
41
42 let eq_mults = {
43 let mut eq_mults = vec![0usize; vk.inner.params.n_stack];
44 for (sort_idx, (air_idx, vdata)) in
45 preflight.proof_shape.sorted_trace_vdata.iter().enumerate()
46 {
47 if vdata.log_height > vk.inner.params.l_skip {
48 let need_rot = vk.inner.per_air[*air_idx].params.need_rot;
49 let n = vdata.log_height - vk.inner.params.l_skip;
50 eq_mults[n - 1] += proof.batch_constraint_proof.column_openings
51 [sort_idx]
52 .iter()
53 .flatten()
54 .count()
55 / if need_rot { 2 } else { 1 };
56 }
57 }
58 eq_mults
59 };
60
61 let u_mults = {
62 let mut u_mults = vec![0usize; vk.inner.params.n_stack];
63 let stacked_slices =
64 get_stacked_slice_data(vk, &preflight.proof_shape.sorted_trace_vdata);
65
66 let mut b_value_set = HashSet::<(usize, usize)>::new();
67 for slice in stacked_slices {
68 let n_lift = slice.n.max(0) as usize;
69 let b_value = slice.row_idx >> (n_lift + vk.inner.params.l_skip);
70 let total_num_bits = vk.inner.params.n_stack - n_lift;
71
72 for num_bits in (1..=total_num_bits).rev() {
73 let shifted_b_value = b_value >> (total_num_bits - num_bits);
74 if b_value_set.insert((shifted_b_value, num_bits)) {
75 u_mults[vk.inner.params.n_stack - num_bits] += 1;
76 } else {
77 break;
78 }
79 }
80 }
81 u_mults
82 };
83
84 let (eq_prism_base, eq_cube_base, rot_cube_base) = {
85 let l_skip = vk.inner.params.l_skip;
86 let omega = F::two_adic_generator(l_skip);
87 let u = preflight.stacking.sumcheck_rnd[0];
88 let r = preflight.batch_constraint.sumcheck_rnd[0];
89
90 let eq_prism_base = eval_eq_uni(l_skip, u, r);
91 let eq_cube_base = eval_eq_uni(l_skip, u, r * omega);
92 let rot_cube_base =
93 eval_eq_uni_at_one(l_skip, u) * eval_eq_uni_at_one(l_skip, r * omega);
94 (eq_prism_base, eq_cube_base, rot_cube_base)
95 };
96
97 let num_rows = sumcheck_rounds.len();
98 let proof_idx_value = F::from_usize(proof_idx);
99
100 let mut trace = vec![F::ZERO; num_rows * width];
101
102 let u = &preflight.stacking.sumcheck_rnd[1..];
103 let batch_sumcheck_randomness = preflight.batch_constraint_sumcheck_randomness();
104 let r = &batch_sumcheck_randomness[1..];
105
106 let initial_tidx = preflight.stacking.intermediate_tidx[1];
107
108 let mut s_eval_at_u = preflight.stacking.univariate_poly_rand_eval;
109
110 let mut eq_cube = EF::ONE;
111 let mut r_not_u_prod = EF::ONE;
112 let mut rot_cube_minus_prod = EF::ZERO;
113
114 for (round, (sumcheck_round, chunk, &u_round)) in
115 izip!(sumcheck_rounds.iter(), trace.chunks_mut(width), u.iter()).enumerate()
116 {
117 let cols: &mut SumcheckRoundsCols<F> = chunk.borrow_mut();
118
119 let s_eval_at_0 = s_eval_at_u - sumcheck_round[0];
120 s_eval_at_u = interpolate_quadratic_at_012(
121 &[s_eval_at_0, sumcheck_round[0], sumcheck_round[1]],
122 u_round,
123 );
124
125 cols.proof_idx = proof_idx_value;
126 cols.is_valid = F::ONE;
127 cols.is_first = F::from_bool(round == 0);
128 cols.is_last = F::from_bool(round + 1 == num_rows);
129
130 cols.round = F::from_usize(round + 1);
131 cols.tidx = F::from_usize(initial_tidx + (3 * D_EF * round));
132
133 cols.s_eval_at_0
134 .copy_from_slice(s_eval_at_0.as_basis_coefficients_slice());
135 cols.s_eval_at_1
136 .copy_from_slice(sumcheck_round[0].as_basis_coefficients_slice());
137 cols.s_eval_at_2
138 .copy_from_slice(sumcheck_round[1].as_basis_coefficients_slice());
139 cols.s_eval_at_u
140 .copy_from_slice(s_eval_at_u.as_basis_coefficients_slice());
141
142 cols.u_round
143 .copy_from_slice(u_round.as_basis_coefficients_slice());
144 let r_round = if round < r.len() {
145 cols.r_round = r[round].challenge;
146 cols.has_r = F::ONE;
147 EF::from_basis_coefficients_iter(r[round].challenge.into_iter()).unwrap()
148 } else {
149 EF::ZERO
150 };
151 cols.u_mult = F::from_usize(u_mults[round]);
152
153 cols.eq_prism_base
154 .copy_from_slice(eq_prism_base.as_basis_coefficients_slice());
155 cols.eq_cube_base
156 .copy_from_slice(eq_cube_base.as_basis_coefficients_slice());
157 cols.rot_cube_base
158 .copy_from_slice(rot_cube_base.as_basis_coefficients_slice());
159
160 let u_not_r = u_round * (EF::ONE - r_round);
161 let r_not_u = r_round * (EF::ONE - u_round);
162 let next_eq_term = EF::ONE - (u_not_r + r_not_u);
163 eq_cube *= next_eq_term;
164 cols.eq_cube
165 .copy_from_slice(eq_cube.as_basis_coefficients_slice());
166
167 rot_cube_minus_prod =
168 (rot_cube_minus_prod * next_eq_term) + u_not_r * r_not_u_prod;
169 r_not_u_prod *= r_not_u;
170 cols.r_not_u_prod
171 .copy_from_slice(r_not_u_prod.as_basis_coefficients_slice());
172 cols.rot_cube_minus_prod
173 .copy_from_slice(rot_cube_minus_prod.as_basis_coefficients_slice());
174
175 cols.eq_rot_mult = F::from_usize(eq_mults[round]);
176 }
177
178 (trace, num_rows)
179 })
180 .collect::<Vec<_>>();
181
182 let num_valid_rows = traces.iter().map(|(_trace, num_rows)| *num_rows).sum();
183 let height = if let Some(height) = required_height {
184 if height < num_valid_rows {
185 return None;
186 }
187 height
188 } else {
189 num_valid_rows.next_power_of_two()
190 };
191
192 let mut combined_trace = Vec::with_capacity(height * width);
193 for (trace, _num_rows) in traces {
194 combined_trace.extend(trace);
195 }
196
197 let padding_proof_idx = F::from_usize(proofs.len());
198 combined_trace.resize(height * width, F::ZERO);
199 let mut chunks = combined_trace[num_valid_rows * width..]
200 .chunks_mut(width)
201 .peekable();
202
203 while let Some(chunk) = chunks.next() {
204 let cols: &mut SumcheckRoundsCols<F> = chunk.borrow_mut();
205 cols.proof_idx = padding_proof_idx;
206 if chunks.peek().is_none() {
207 cols.is_last = F::ONE;
208 }
209 }
210
211 Some(RowMajorMatrix::new(combined_trace, width))
212 }
213}