openvm_recursion_circuit/stacking/sumcheck/
trace.rs

1use 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}