openvm_recursion_circuit/gkr/sumcheck/
trace.rs1use 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 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 let rows_per_proof: Vec<usize> = gkr_sumcheck_records
80 .iter()
81 .map(|record| record.total_rounds().max(1))
82 .collect();
83
84 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 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 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}