openvm_recursion_circuit/gkr/input/
trace.rs1use 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 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 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 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}