openvm_recursion_circuit/stacking/univariate/
trace.rs

1use std::borrow::BorrowMut;
2
3use itertools::{izip, Itertools};
4use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, EF, F};
5use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
6use p3_matrix::dense::RowMajorMatrix;
7use p3_maybe_rayon::prelude::*;
8
9use crate::{
10    stacking::univariate::air::UnivariateRoundCols,
11    tracegen::{RowMajorChip, StandardTracegenCtx},
12};
13
14pub struct UnivariateRoundTraceGenerator;
15
16impl RowMajorChip<F> for UnivariateRoundTraceGenerator {
17    type Ctx<'a> = StandardTracegenCtx<'a>;
18
19    #[tracing::instrument(level = "trace", skip_all)]
20    fn generate_trace(
21        &self,
22        ctx: &Self::Ctx<'_>,
23        required_height: Option<usize>,
24    ) -> Option<RowMajorMatrix<F>> {
25        let vk = ctx.vk;
26        let proofs = ctx.proofs;
27        let preflights = ctx.preflights;
28        debug_assert_eq!(proofs.len(), preflights.len());
29
30        let width = UnivariateRoundCols::<usize>::width();
31
32        let traces = proofs
33            .par_iter()
34            .zip(preflights.par_iter())
35            .enumerate()
36            .map(|(proof_idx, (proof, preflight))| {
37                let coeffs = &proof.stacking_proof.univariate_round_coeffs;
38                let num_rows = coeffs.len();
39                let proof_idx_value = F::from_usize(proof_idx);
40
41                let mut trace = vec![F::ZERO; num_rows * width];
42
43                let u_0 = preflight.stacking.sumcheck_rnd[0];
44                let u_0_pows = u_0.powers().take(num_rows).collect_vec();
45
46                let initial_tidx = preflight.stacking.intermediate_tidx[0];
47
48                let d_card = 1usize << vk.inner.params.l_skip;
49                let mut s_0_sum_over_d = coeffs[0] * F::from_usize(d_card);
50                let mut poly_rand_eval = EF::ZERO;
51
52                for (i, (&coeff, chunk, &u_0_pow)) in
53                    izip!(coeffs.iter(), trace.chunks_mut(width), u_0_pows.iter()).enumerate()
54                {
55                    let cols: &mut UnivariateRoundCols<F> = chunk.borrow_mut();
56                    cols.proof_idx = proof_idx_value;
57                    cols.is_valid = F::ONE;
58                    cols.is_first = F::from_bool(i == 0);
59                    cols.is_last = F::from_bool(i + 1 == num_rows);
60
61                    cols.tidx = F::from_usize(initial_tidx + (D_EF * i));
62                    cols.u_0.copy_from_slice(u_0.as_basis_coefficients_slice());
63                    cols.u_0_pow
64                        .copy_from_slice(u_0_pow.as_basis_coefficients_slice());
65
66                    cols.coeff
67                        .copy_from_slice(coeff.as_basis_coefficients_slice());
68
69                    cols.coeff_idx = F::from_usize(i);
70                    if i == d_card {
71                        s_0_sum_over_d += coeff * F::from_usize(d_card);
72                        cols.coeff_is_d = F::ONE;
73                    }
74                    cols.s_0_sum_over_d
75                        .copy_from_slice(s_0_sum_over_d.as_basis_coefficients_slice());
76
77                    poly_rand_eval += coeff * u_0_pow;
78                    cols.poly_rand_eval
79                        .copy_from_slice(poly_rand_eval.as_basis_coefficients_slice());
80                }
81
82                (trace, num_rows)
83            })
84            .collect::<Vec<_>>();
85
86        let num_valid_rows = traces.iter().map(|(_trace, num_rows)| *num_rows).sum();
87        let height = if let Some(height) = required_height {
88            if height < num_valid_rows {
89                return None;
90            }
91            height
92        } else {
93            num_valid_rows.next_power_of_two()
94        };
95
96        let mut combined_trace = Vec::with_capacity(height * width);
97        for (trace, _num_rows) in traces {
98            combined_trace.extend(trace);
99        }
100
101        let padding_proof_idx = F::from_usize(proofs.len());
102        combined_trace.resize(height * width, F::ZERO);
103        let mut chunks = combined_trace[num_valid_rows * width..]
104            .chunks_mut(width)
105            .peekable();
106
107        while let Some(chunk) = chunks.next() {
108            let cols: &mut UnivariateRoundCols<F> = chunk.borrow_mut();
109            cols.proof_idx = padding_proof_idx;
110            if chunks.peek().is_none() {
111                cols.is_last = F::ONE;
112            }
113        }
114
115        Some(RowMajorMatrix::new(combined_trace, width))
116    }
117}