openvm_recursion_circuit/gkr/layer/
trace.rs1use core::borrow::BorrowMut;
2
3use openvm_stark_backend::p3_maybe_rayon::prelude::*;
4use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, EF, F};
5use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
6use p3_matrix::dense::RowMajorMatrix;
7
8use super::{air::reduce_to_single_evaluation, GkrLayerCols};
9use crate::tracegen::RowMajorChip;
10
11#[derive(Debug, Clone, Default)]
13pub struct GkrLayerRecord {
14 pub tidx: usize,
15 pub layer_claims: Vec<[EF; 4]>,
16 pub lambdas: Vec<EF>,
17 pub eq_at_r_primes: Vec<EF>,
18}
19
20impl GkrLayerRecord {
21 #[inline]
22 fn layer_count(&self) -> usize {
23 self.layer_claims.len()
24 }
25
26 #[inline]
27 fn lambda_at(&self, layer_idx: usize) -> EF {
28 layer_idx
29 .checked_sub(1)
30 .and_then(|idx| self.lambdas.get(idx))
31 .copied()
32 .unwrap_or(EF::ZERO)
33 }
34
35 #[inline]
36 fn eq_at(&self, layer_idx: usize) -> EF {
37 layer_idx
38 .checked_sub(1)
39 .and_then(|idx| self.eq_at_r_primes.get(idx))
40 .copied()
41 .unwrap_or(EF::ZERO)
42 }
43
44 #[inline]
45 fn layer_tidx(&self, layer_idx: usize) -> usize {
46 if layer_idx == 0 {
47 self.tidx
48 } else {
49 let j = layer_idx;
50 self.tidx + D_EF * (2 * j * j + 4 * j - 1)
51 }
52 }
53}
54
55pub struct GkrLayerTraceGenerator;
56
57impl RowMajorChip<F> for GkrLayerTraceGenerator {
58 type Ctx<'a> = (&'a [GkrLayerRecord], &'a [Vec<EF>], &'a [EF]);
60
61 #[tracing::instrument(level = "trace", skip_all)]
62 fn generate_trace(
63 &self,
64 ctx: &Self::Ctx<'_>,
65 required_height: Option<usize>,
66 ) -> Option<RowMajorMatrix<F>> {
67 let (gkr_layer_records, mus, q0_claims) = ctx;
68 debug_assert_eq!(gkr_layer_records.len(), mus.len());
69 debug_assert_eq!(gkr_layer_records.len(), q0_claims.len());
70
71 let width = GkrLayerCols::<F>::width();
72
73 let rows_per_proof: Vec<usize> = gkr_layer_records
75 .iter()
76 .map(|record| record.layer_claims.len().max(1))
77 .collect();
78
79 let num_valid_rows: usize = rows_per_proof.iter().sum();
81 let height = if let Some(height) = required_height {
82 if height < num_valid_rows {
83 return None;
84 }
85 height
86 } else {
87 num_valid_rows.next_power_of_two()
88 };
89 let mut trace = vec![F::ZERO; height * width];
90
91 let (data_slice, _) = trace.split_at_mut(num_valid_rows * width);
93 let mut trace_slices: Vec<&mut [F]> = Vec::with_capacity(rows_per_proof.len());
94 let mut remaining = data_slice;
95
96 for &num_rows in &rows_per_proof {
97 let chunk_size = num_rows * width;
98 let (chunk, rest) = remaining.split_at_mut(chunk_size);
99 trace_slices.push(chunk);
100 remaining = rest;
101 }
102
103 trace_slices
105 .par_iter_mut()
106 .zip(
107 gkr_layer_records
108 .par_iter()
109 .zip(mus.par_iter())
110 .zip(q0_claims.par_iter()),
111 )
112 .enumerate()
113 .for_each(
114 |(proof_idx, (proof_trace, ((record, mus_for_proof), q0_claim)))| {
115 let mus_for_proof = mus_for_proof.as_slice();
116 let q0_claim = *q0_claim;
117
118 if record.layer_claims.is_empty() {
119 debug_assert_eq!(proof_trace.len(), width);
120 let row_data = &mut proof_trace[..width];
121 let cols: &mut GkrLayerCols<F> = row_data.borrow_mut();
122 cols.is_enabled = F::ONE;
123 cols.proof_idx = F::from_usize(proof_idx);
124 cols.is_first = F::ONE;
125 cols.is_dummy = F::ONE;
126 cols.sumcheck_claim_in = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
127 cols.q_xi_0 = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
128 cols.q_xi_1 = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
129 cols.denom_claim = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
130 return;
131 }
132
133 let layer_count = record.layer_count();
134 let mut prev_layer_eval: Option<(EF, EF)> = None;
135
136 proof_trace
137 .chunks_mut(width)
138 .take(layer_count)
139 .enumerate()
140 .for_each(|(layer_idx, row_data)| {
141 let cols: &mut GkrLayerCols<F> = row_data.borrow_mut();
142 cols.proof_idx = F::from_usize(proof_idx);
143 cols.is_enabled = F::ONE;
144 cols.is_first = F::from_bool(layer_idx == 0);
145 cols.layer_idx = F::from_usize(layer_idx);
146 cols.tidx = F::from_usize(record.layer_tidx(layer_idx));
147
148 let lambda = record.lambda_at(layer_idx);
149 let eq_at_r_prime = record.eq_at(layer_idx);
150
151 cols.lambda = lambda.as_basis_coefficients_slice().try_into().unwrap();
152 cols.eq_at_r_prime = eq_at_r_prime
153 .as_basis_coefficients_slice()
154 .try_into()
155 .unwrap();
156
157 let claims = &record.layer_claims[layer_idx];
158 let mu = mus_for_proof[layer_idx];
159
160 cols.p_xi_0 =
161 claims[0].as_basis_coefficients_slice().try_into().unwrap();
162 cols.q_xi_0 =
163 claims[1].as_basis_coefficients_slice().try_into().unwrap();
164 cols.p_xi_1 =
165 claims[2].as_basis_coefficients_slice().try_into().unwrap();
166 cols.q_xi_1 =
167 claims[3].as_basis_coefficients_slice().try_into().unwrap();
168
169 cols.mu = mu.as_basis_coefficients_slice().try_into().unwrap();
170
171 let sumcheck_claim_in = prev_layer_eval
172 .map(|(numer_prev, denom_prev)| numer_prev + lambda * denom_prev)
173 .unwrap_or(q0_claim);
174 cols.sumcheck_claim_in = sumcheck_claim_in
175 .as_basis_coefficients_slice()
176 .try_into()
177 .unwrap();
178
179 let (numer_base, denom_base): ([F; D_EF], [F; D_EF]) =
180 reduce_to_single_evaluation::<F, F>(
181 claims[0].as_basis_coefficients_slice().try_into().unwrap(),
182 claims[2].as_basis_coefficients_slice().try_into().unwrap(),
183 claims[1].as_basis_coefficients_slice().try_into().unwrap(),
184 claims[3].as_basis_coefficients_slice().try_into().unwrap(),
185 mu.as_basis_coefficients_slice().try_into().unwrap(),
186 );
187 cols.numer_claim = numer_base;
188 cols.denom_claim = denom_base;
189
190 let numer = claims[0] * (EF::ONE - mu) + claims[2] * mu;
191 let denom = claims[1] * (EF::ONE - mu) + claims[3] * mu;
192 prev_layer_eval = Some((numer, denom));
193 });
194 },
195 );
196
197 Some(RowMajorMatrix::new(trace, width))
198 }
199}