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