openvm_recursion_circuit/gkr/sumcheck/
trace.rs

1use core::borrow::BorrowMut;
2
3use openvm_stark_backend::{p3_maybe_rayon::prelude::*, poly_common::interpolate_cubic_at_0123};
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::GkrLayerSumcheckCols;
9use crate::tracegen::RowMajorChip;
10
11#[derive(Default, Debug, Clone)]
12pub struct GkrSumcheckRecord {
13    pub tidx: usize,
14    pub evals: Vec<[EF; 3]>,
15    pub ris: Vec<EF>,
16    pub claims: Vec<EF>,
17}
18
19impl GkrSumcheckRecord {
20    #[inline]
21    pub fn num_layers(&self) -> usize {
22        self.claims.len()
23    }
24
25    #[inline]
26    pub fn total_rounds(&self) -> usize {
27        let layers = self.num_layers();
28        layers * (layers + 1) / 2
29    }
30
31    #[inline]
32    fn layer_start_index(layer_idx: usize) -> usize {
33        layer_idx * (layer_idx + 1) / 2
34    }
35
36    #[inline]
37    fn layer_rounds(layer_idx: usize) -> usize {
38        layer_idx + 1
39    }
40
41    #[inline]
42    fn derive_tidx(&self, layer_idx: usize, round_in_layer: usize) -> usize {
43        let rounds_before_layer = Self::layer_start_index(layer_idx);
44        self.tidx + 4 * D_EF * (rounds_before_layer + round_in_layer) + 6 * D_EF * layer_idx
45    }
46
47    #[inline]
48    fn prev_challenge(layer_idx: usize, round_in_layer: usize, mus: &[EF], ris: &[EF]) -> EF {
49        if round_in_layer == 0 {
50            mus[layer_idx]
51        } else {
52            let prev_layer = layer_idx
53                .checked_sub(1)
54                .expect("round_in_layer > 0 only occurs for non-root layers");
55            let offset = Self::layer_start_index(prev_layer) + (round_in_layer - 1);
56            ris[offset]
57        }
58    }
59}
60
61pub struct GkrSumcheckTraceGenerator;
62
63impl RowMajorChip<F> for GkrSumcheckTraceGenerator {
64    // (gkr_sumcheck_records, mus)
65    type Ctx<'a> = (&'a [GkrSumcheckRecord], &'a [Vec<EF>]);
66
67    #[tracing::instrument(level = "trace", skip_all)]
68    fn generate_trace(
69        &self,
70        ctx: &Self::Ctx<'_>,
71        required_height: Option<usize>,
72    ) -> Option<RowMajorMatrix<F>> {
73        let (gkr_sumcheck_records, mus) = ctx;
74        debug_assert_eq!(gkr_sumcheck_records.len(), mus.len());
75
76        let width = GkrLayerSumcheckCols::<F>::width();
77
78        // Calculate rows per proof
79        let rows_per_proof: Vec<usize> = gkr_sumcheck_records
80            .iter()
81            .map(|record| record.total_rounds().max(1))
82            .collect();
83
84        // Calculate total rows
85        let num_valid_rows: usize = rows_per_proof.iter().sum();
86        let height = if let Some(height) = required_height {
87            if height < num_valid_rows {
88                return None;
89            }
90            height
91        } else {
92            num_valid_rows.next_power_of_two()
93        };
94        let mut trace = vec![F::ZERO; height * width];
95
96        // Split trace into chunks for each proof
97        let (data_slice, _) = trace.split_at_mut(num_valid_rows * width);
98        let mut trace_slices: Vec<&mut [F]> = Vec::with_capacity(rows_per_proof.len());
99        let mut remaining = data_slice;
100
101        for &num_rows in &rows_per_proof {
102            let chunk_size = num_rows * width;
103            let (chunk, rest) = remaining.split_at_mut(chunk_size);
104            trace_slices.push(chunk);
105            remaining = rest;
106        }
107
108        // Process each proof in parallel
109        trace_slices
110            .par_iter_mut()
111            .zip(gkr_sumcheck_records.par_iter().zip(mus.par_iter()))
112            .enumerate()
113            .for_each(|(proof_idx, (proof_trace, (record, mus_for_proof)))| {
114                let mus_for_proof = mus_for_proof.as_slice();
115                let total_rounds = record.total_rounds();
116                let num_layers = record.num_layers();
117
118                debug_assert_eq!(record.ris.len(), total_rounds);
119                debug_assert_eq!(record.evals.len(), total_rounds);
120                debug_assert!(mus_for_proof.len() >= num_layers);
121
122                if total_rounds == 0 {
123                    debug_assert_eq!(proof_trace.len(), width);
124                    let row_data = &mut proof_trace[..width];
125                    let cols: &mut GkrLayerSumcheckCols<F> = row_data.borrow_mut();
126                    cols.is_enabled = F::ONE;
127                    cols.tidx = F::from_usize(D_EF);
128                    cols.proof_idx = F::from_usize(proof_idx);
129                    cols.layer_idx = F::ONE;
130                    cols.is_first_round = F::ONE;
131                    cols.is_proof_start = F::ONE;
132                    cols.is_last_layer = F::ONE;
133                    cols.is_dummy = F::ONE;
134                    cols.eq_in = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
135                    cols.eq_out = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
136                    cols.claim_in = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
137                    cols.claim_out = [F::ONE, F::ZERO, F::ZERO, F::ZERO];
138                    return;
139                }
140
141                let mut global_round_idx = 0usize;
142                let mut row_iter = proof_trace.chunks_mut(width);
143
144                for layer_idx in 0..num_layers {
145                    let layer_rounds = GkrSumcheckRecord::layer_rounds(layer_idx);
146                    let layer_idx_value = layer_idx + 1;
147                    let is_last_layer = layer_idx == num_layers.saturating_sub(1);
148
149                    let mut claim = record.claims[layer_idx];
150                    let mut eq = EF::ONE;
151
152                    for round_in_layer in 0..layer_rounds {
153                        let challenge = record.ris[global_round_idx];
154                        let evals = record.evals[global_round_idx];
155                        let prev_challenge = GkrSumcheckRecord::prev_challenge(
156                            layer_idx,
157                            round_in_layer,
158                            mus_for_proof,
159                            &record.ris,
160                        );
161
162                        let prev_challenge_base: [F; D_EF] = prev_challenge
163                            .as_basis_coefficients_slice()
164                            .try_into()
165                            .unwrap();
166                        let challenge_base: [F; D_EF] =
167                            challenge.as_basis_coefficients_slice().try_into().unwrap();
168
169                        let eval1_base: [F; D_EF] =
170                            evals[0].as_basis_coefficients_slice().try_into().unwrap();
171                        let eval2_base: [F; D_EF] =
172                            evals[1].as_basis_coefficients_slice().try_into().unwrap();
173                        let eval3_base: [F; D_EF] =
174                            evals[2].as_basis_coefficients_slice().try_into().unwrap();
175
176                        let claim_in_base: [F; D_EF] =
177                            claim.as_basis_coefficients_slice().try_into().unwrap();
178                        let eq_in_base: [F; D_EF] =
179                            eq.as_basis_coefficients_slice().try_into().unwrap();
180
181                        let ev0 = claim - evals[0];
182                        let evals_full = [ev0, evals[0], evals[1], evals[2]];
183                        let claim_out = interpolate_cubic_at_0123(&evals_full, challenge);
184                        let eq_factor = prev_challenge * challenge
185                            + (EF::ONE - prev_challenge) * (EF::ONE - challenge);
186                        let eq_out = eq * eq_factor;
187
188                        let claim_out_base: [F; D_EF] =
189                            claim_out.as_basis_coefficients_slice().try_into().unwrap();
190                        let eq_out_base: [F; D_EF] =
191                            eq_out.as_basis_coefficients_slice().try_into().unwrap();
192
193                        let cols: &mut GkrLayerSumcheckCols<F> =
194                            row_iter.next().unwrap().borrow_mut();
195                        cols.is_enabled = F::ONE;
196                        cols.proof_idx = F::from_usize(proof_idx);
197
198                        cols.layer_idx = F::from_usize(layer_idx_value);
199                        cols.is_last_layer = F::from_bool(is_last_layer);
200
201                        cols.round = F::from_usize(round_in_layer);
202                        cols.is_first_round = F::from_bool(round_in_layer == 0);
203                        cols.is_proof_start =
204                            F::from_bool(layer_idx_value == 1 && round_in_layer == 0);
205
206                        let tidx = record.derive_tidx(layer_idx, round_in_layer);
207                        cols.tidx = F::from_usize(tidx);
208
209                        cols.ev1 = eval1_base;
210                        cols.ev2 = eval2_base;
211                        cols.ev3 = eval3_base;
212
213                        cols.prev_challenge = prev_challenge_base;
214                        cols.challenge = challenge_base;
215
216                        cols.claim_in = claim_in_base;
217                        cols.claim_out = claim_out_base;
218
219                        cols.eq_in = eq_in_base;
220                        cols.eq_out = eq_out_base;
221
222                        claim = claim_out;
223                        eq = eq_out;
224                        global_round_idx += 1;
225                    }
226                }
227
228                debug_assert_eq!(global_round_idx, total_rounds);
229            });
230
231        Some(RowMajorMatrix::new(trace, width))
232    }
233}