openvm_recursion_circuit/gkr/input/
trace.rs

1use core::borrow::BorrowMut;
2
3use openvm_circuit_primitives::{is_zero::IsZeroSubAir, TraceSubRowGenerator};
4use openvm_stark_backend::p3_maybe_rayon::prelude::*;
5use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
6use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
7use p3_matrix::dense::RowMajorMatrix;
8
9use super::GkrInputCols;
10use crate::tracegen::RowMajorChip;
11
12#[derive(Debug, Clone, Default)]
13pub struct GkrInputRecord {
14    pub tidx: usize,
15    pub n_logup: usize,
16    pub n_max: usize,
17    pub logup_pow_witness: F,
18    pub logup_pow_sample: F,
19    pub alpha_logup: EF,
20    pub input_layer_claim: [EF; 2],
21}
22
23pub struct GkrInputTraceGenerator;
24
25impl RowMajorChip<F> for GkrInputTraceGenerator {
26    // (gkr_input_records, q0_claims)
27    type Ctx<'a> = (&'a [GkrInputRecord], &'a [EF]);
28
29    #[tracing::instrument(level = "trace", skip_all)]
30    fn generate_trace(
31        &self,
32        ctx: &Self::Ctx<'_>,
33        required_height: Option<usize>,
34    ) -> Option<RowMajorMatrix<F>> {
35        let (gkr_input_records, q0_claims) = ctx;
36        debug_assert_eq!(gkr_input_records.len(), q0_claims.len());
37
38        let width = GkrInputCols::<F>::width();
39
40        // Each record generates exactly 1 row
41        let num_valid_rows = gkr_input_records.len();
42        let height = if let Some(height) = required_height {
43            if height < num_valid_rows {
44                return None;
45            }
46            height
47        } else {
48            num_valid_rows.next_power_of_two()
49        };
50        let mut trace = vec![F::ZERO; height * width];
51
52        let (data_slice, _) = trace.split_at_mut(num_valid_rows * width);
53
54        // Process each proof row
55        data_slice
56            .par_chunks_mut(width)
57            .zip(gkr_input_records.par_iter().zip(q0_claims.par_iter()))
58            .enumerate()
59            .for_each(|(proof_idx, (row_data, (record, q0_claim)))| {
60                let cols: &mut GkrInputCols<F> = row_data.borrow_mut();
61
62                cols.is_enabled = F::ONE;
63                cols.proof_idx = F::from_usize(proof_idx);
64
65                cols.tidx = F::from_usize(record.tidx);
66
67                cols.n_logup = F::from_usize(record.n_logup);
68                cols.n_max = F::from_usize(record.n_max);
69                cols.is_n_max_greater_than_n_logup = F::from_bool(record.n_max > record.n_logup);
70
71                IsZeroSubAir.generate_subrow(
72                    cols.n_logup,
73                    (&mut cols.is_n_logup_zero_aux.inv, &mut cols.is_n_logup_zero),
74                );
75
76                cols.logup_pow_witness = record.logup_pow_witness;
77                cols.logup_pow_sample = record.logup_pow_sample;
78
79                cols.q0_claim = q0_claim.as_basis_coefficients_slice().try_into().unwrap();
80                cols.alpha_logup = record
81                    .alpha_logup
82                    .as_basis_coefficients_slice()
83                    .try_into()
84                    .unwrap();
85                cols.input_layer_claim = [
86                    record.input_layer_claim[0]
87                        .as_basis_coefficients_slice()
88                        .try_into()
89                        .unwrap(),
90                    record.input_layer_claim[1]
91                        .as_basis_coefficients_slice()
92                        .try_into()
93                        .unwrap(),
94                ];
95            });
96
97        Some(RowMajorMatrix::new(trace, width))
98    }
99}