openvm_recursion_circuit/batch_constraint/eq_airs/eq_uni/
trace.rs

1use std::borrow::BorrowMut;
2
3use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
4use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
5use p3_matrix::dense::RowMajorMatrix;
6
7use crate::{
8    batch_constraint::eq_airs::eq_uni::air::EqUniCols,
9    tracegen::{RowMajorChip, StandardTracegenCtx},
10};
11
12pub struct EqUniTraceGenerator;
13
14impl RowMajorChip<F> for EqUniTraceGenerator {
15    type Ctx<'a> = StandardTracegenCtx<'a>;
16
17    #[tracing::instrument(level = "trace", skip_all)]
18    fn generate_trace(
19        &self,
20        ctx: &Self::Ctx<'_>,
21        required_height: Option<usize>,
22    ) -> Option<RowMajorMatrix<F>> {
23        let vk = ctx.vk;
24        let preflights = ctx.preflights;
25        let width = EqUniCols::<F>::width();
26        let l_skip = vk.inner.params.l_skip;
27        let one_height = l_skip + 1;
28        let total_height = one_height * preflights.len();
29        let padding_height = if let Some(height) = required_height {
30            if height < total_height {
31                return None;
32            }
33            height
34        } else {
35            total_height.next_power_of_two()
36        };
37        let mut trace = vec![F::ZERO; padding_height * width];
38
39        for (pidx, preflight) in preflights.iter().enumerate() {
40            let mut x = preflight.batch_constraint.xi[0];
41            let mut y = preflight.batch_constraint.sumcheck_rnd[0];
42            let mut res = EF::ONE;
43            trace[pidx * one_height * width..(pidx + 1) * one_height * width]
44                .chunks_exact_mut(width)
45                .enumerate()
46                .for_each(|(i, chunk)| {
47                    let cols: &mut EqUniCols<_> = chunk.borrow_mut();
48                    cols.is_valid = F::ONE;
49                    cols.proof_idx = F::from_usize(pidx);
50                    cols.is_first = F::from_bool(i == 0);
51
52                    cols.idx = F::from_usize(i);
53                    cols.x.copy_from_slice(x.as_basis_coefficients_slice());
54                    cols.y.copy_from_slice(y.as_basis_coefficients_slice());
55                    cols.res.copy_from_slice(res.as_basis_coefficients_slice());
56
57                    res = (x + y) * res + (EF::ONE - x) * (EF::ONE - y);
58                    x *= x;
59                    y *= y;
60                });
61        }
62
63        trace[total_height * width..]
64            .chunks_mut(width)
65            .enumerate()
66            .for_each(|(i, chunk)| {
67                let cols: &mut EqUniCols<F> = chunk.borrow_mut();
68                cols.proof_idx = F::from_usize(preflights.len() + i);
69            });
70
71        Some(RowMajorMatrix::new(trace, width))
72    }
73}