openvm_recursion_circuit/gkr/layer/
trace.rs

1use 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/// Minimal record for parallel gkr layer trace generation
12#[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    // (gkr_layer_records, mus, q0_claims)
59    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        // Calculate rows per proof (each record has layer_claims.len() rows)
74        let rows_per_proof: Vec<usize> = gkr_layer_records
75            .iter()
76            .map(|record| record.layer_claims.len().max(1))
77            .collect();
78
79        // Calculate total rows
80        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        // Split trace into chunks for each proof and process in parallel
92        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        // Process each proof in parallel
104        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}