openvm_recursion_circuit/stacking/opening/
trace.rs

1use std::{borrow::BorrowMut, iter::zip};
2
3use itertools::{izip, Itertools};
4use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, EF, F};
5use p3_field::{BasedVectorSpace, Field, PrimeCharacteristicRing};
6use p3_matrix::dense::RowMajorMatrix;
7
8use crate::{
9    stacking::{
10        opening::air::OpeningClaimsCols,
11        utils::{
12            compute_coefficients, get_stacked_slice_data, sorted_column_claims, ColumnOpeningPair,
13        },
14    },
15    tracegen::{RowMajorChip, StandardTracegenCtx},
16};
17
18pub struct OpeningClaimsTraceGenerator;
19
20impl RowMajorChip<F> for OpeningClaimsTraceGenerator {
21    type Ctx<'a> = StandardTracegenCtx<'a>;
22
23    #[tracing::instrument(level = "trace", skip_all)]
24    fn generate_trace(
25        &self,
26        ctx: &Self::Ctx<'_>,
27        required_height: Option<usize>,
28    ) -> Option<RowMajorMatrix<F>> {
29        let vk = ctx.vk;
30        let proofs = ctx.proofs;
31        let preflights = ctx.preflights;
32        debug_assert_eq!(proofs.len(), preflights.len());
33
34        let width = OpeningClaimsCols::<usize>::width();
35        let column_claims = zip(proofs.iter(), preflights.iter())
36            .map(|(proof, preflight)| {
37                sorted_column_claims(vk, proof, &preflight.proof_shape.sorted_trace_vdata)
38            })
39            .collect_vec();
40        let minimum_height: usize = column_claims.iter().map(|c| c.len()).sum();
41        let height = if let Some(height) = required_height {
42            if height < minimum_height {
43                return None;
44            }
45            height
46        } else {
47            minimum_height.next_power_of_two()
48        };
49
50        let mut trace = vec![F::ZERO; height * width];
51        let mut chunks = trace.chunks_mut(width);
52
53        for (proof_idx, (proof, preflight, claims)) in
54            izip!(proofs, preflights, column_claims).enumerate()
55        {
56            let stacked_slices =
57                get_stacked_slice_data(vk, &preflight.proof_shape.sorted_trace_vdata);
58
59            let (_, per_slice) = compute_coefficients(
60                proof,
61                &stacked_slices,
62                &preflight.stacking.sumcheck_rnd,
63                &preflight.batch_constraint.sumcheck_rnd,
64                &preflight.stacking.lambda,
65                vk.inner.params.l_skip,
66                vk.inner.params.n_stack,
67            );
68
69            let num_rows = claims.len();
70            let proof_idx_value = F::from_usize(proof_idx);
71
72            let mut lambda_pows = preflight.stacking.lambda.square().powers().take(num_rows);
73            let mut stacking_claim_coefficient = EF::ZERO;
74            let mut s_0 = EF::ZERO;
75
76            let last_main_idx = claims
77                .iter()
78                .enumerate()
79                .skip(1)
80                .find_map(|(i, claim)| {
81                    if claim.part_idx != 0 {
82                        Some(i - 1)
83                    } else {
84                        None
85                    }
86                })
87                .unwrap_or(num_rows - 1);
88
89            for (row_idx, (claim, slice, (eq_in, k_rot_in, eq_bits))) in
90                izip!(claims, stacked_slices, per_slice,).enumerate()
91            {
92                let chunk = chunks.next().unwrap();
93                let ColumnOpeningPair {
94                    sort_idx,
95                    part_idx,
96                    col_idx,
97                    col_claim,
98                    rot_claim,
99                } = claim;
100                let cols: &mut OpeningClaimsCols<F> = chunk.borrow_mut();
101                cols.proof_idx = proof_idx_value;
102                cols.is_valid = F::ONE;
103                cols.is_first = F::from_bool(row_idx == 0);
104                cols.is_last = F::from_bool(row_idx + 1 == num_rows);
105
106                cols.sort_idx = F::from_usize(sort_idx);
107                cols.part_idx = F::from_usize(part_idx);
108                cols.col_idx = F::from_usize(col_idx);
109                cols.col_claim
110                    .copy_from_slice(col_claim.as_basis_coefficients_slice());
111                cols.rot_claim
112                    .copy_from_slice(rot_claim.as_basis_coefficients_slice());
113                cols.need_rot = F::from_bool(slice.need_rot);
114
115                cols.is_main = F::from_bool(part_idx == 0);
116                cols.is_transition_main =
117                    F::from_bool(row_idx + 1 != num_rows && row_idx != last_main_idx);
118
119                let n_lift = slice.n.max(0) as usize;
120                cols.hypercube_dim = if slice.n.is_positive() {
121                    F::from_usize(n_lift)
122                } else {
123                    -F::from_usize(slice.n.unsigned_abs())
124                };
125                cols.log_lifted_height = F::from_usize(n_lift + vk.inner.params.l_skip);
126                cols.lifted_height = F::from_usize(1 << (n_lift + vk.inner.params.l_skip));
127                cols.lifted_height_inv = cols.lifted_height.inverse();
128
129                cols.tidx = F::from_usize(
130                    preflight.batch_constraint.tidx_before_column_openings + 2 * row_idx * D_EF,
131                );
132                cols.lambda
133                    .copy_from_slice(preflight.stacking.lambda.as_basis_coefficients_slice());
134
135                let lambda_pow = lambda_pows.next().unwrap();
136                cols.lambda_pow
137                    .copy_from_slice(lambda_pow.as_basis_coefficients_slice());
138
139                cols.commit_idx = F::from_usize(slice.commit_idx);
140                cols.stacked_col_idx = F::from_usize(slice.col_idx);
141                cols.row_idx = F::from_usize(slice.row_idx);
142                cols.is_last_for_claim = F::from_bool(slice.is_last_for_claim);
143
144                cols.eq_in
145                    .copy_from_slice(eq_in.as_basis_coefficients_slice());
146                cols.k_rot_in
147                    .copy_from_slice(k_rot_in.as_basis_coefficients_slice());
148                if slice.need_rot {
149                    cols.k_rot_in_when_needed
150                        .copy_from_slice(k_rot_in.as_basis_coefficients_slice());
151                }
152                cols.eq_bits
153                    .copy_from_slice(eq_bits.as_basis_coefficients_slice());
154
155                let lambda_pow_eq_bits = lambda_pow * eq_bits;
156                cols.lambda_pow_eq_bits
157                    .copy_from_slice(lambda_pow_eq_bits.as_basis_coefficients_slice());
158
159                let k_rot_term = if slice.need_rot { k_rot_in } else { EF::ZERO };
160                stacking_claim_coefficient +=
161                    lambda_pow_eq_bits * (eq_in + preflight.stacking.lambda * k_rot_term);
162                cols.stacking_claim_coefficient
163                    .copy_from_slice(stacking_claim_coefficient.as_basis_coefficients_slice());
164                if slice.is_last_for_claim {
165                    stacking_claim_coefficient = EF::ZERO;
166                }
167
168                s_0 += lambda_pow * (col_claim + preflight.stacking.lambda * rot_claim);
169                cols.s_0.copy_from_slice(s_0.as_basis_coefficients_slice());
170            }
171        }
172
173        let padding_proof_idx = F::from_usize(proofs.len());
174        let mut chunks = chunks.peekable();
175
176        while let Some(chunk) = chunks.next() {
177            let cols: &mut OpeningClaimsCols<F> = chunk.borrow_mut();
178            cols.proof_idx = padding_proof_idx;
179            if chunks.peek().is_none() {
180                cols.is_last = F::ONE;
181            }
182        }
183
184        Some(RowMajorMatrix::new(trace, width))
185    }
186}
187
188#[cfg(feature = "cuda")]
189pub(crate) mod cuda {
190    use itertools::Itertools;
191    use openvm_cuda_backend::{base::DeviceMatrix, GpuBackend};
192    use openvm_cuda_common::{copy::MemCopyH2D, d_buffer::DeviceBuffer};
193    use openvm_stark_backend::prover::AirProvingContext;
194
195    use super::*;
196    use crate::{
197        stacking::{
198            cuda_abi::{
199                opening_claims_tracegen, opening_claims_tracegen_temp_bytes, ColumnOpeningClaims,
200                OpeningRecordsPerProof,
201            },
202            cuda_tracegen::StackingBlob,
203        },
204        tracegen::{cuda::StandardTracegenGpuCtx, ModuleChip},
205    };
206
207    pub struct OpeningClaimsTraceGeneratorGpu;
208
209    impl ModuleChip<GpuBackend> for OpeningClaimsTraceGeneratorGpu {
210        type Ctx<'a> = (StandardTracegenGpuCtx<'a>, &'a StackingBlob);
211
212        fn generate_proving_ctx(
213            &self,
214            ctx: &Self::Ctx<'_>,
215            required_height: Option<usize>,
216        ) -> Option<AirProvingContext<GpuBackend>> {
217            let child_vk = ctx.0.vk;
218            let proofs_gpu = ctx.0.proofs;
219            let preflights_gpu = ctx.0.preflights;
220            let device_ctx = ctx.0.device_ctx;
221            let blob = ctx.1;
222
223            let mut num_valid_rows = 0;
224            let row_bounds = blob
225                .slice_data
226                .iter()
227                .map(|buf| {
228                    num_valid_rows += buf.len();
229                    num_valid_rows as u32
230                })
231                .collect_vec();
232            let mut last_main_idx_per_proof = Vec::with_capacity(proofs_gpu.len());
233            let claims = proofs_gpu
234                .iter()
235                .zip_eq(preflights_gpu.iter())
236                .map(|(proof, preflight)| {
237                    let claims = sorted_column_claims(
238                        &child_vk.cpu,
239                        &proof.cpu,
240                        &preflight.cpu.proof_shape.sorted_trace_vdata,
241                    )
242                    .into_iter()
243                    .map(|claim| ColumnOpeningClaims {
244                        sort_idx: claim.sort_idx as u32,
245                        part_idx: claim.part_idx as u32,
246                        col_idx: claim.col_idx as u32,
247                        col_claim: claim.col_claim,
248                        rot_claim: claim.rot_claim,
249                    })
250                    .collect_vec();
251                    let last_main_idx = claims
252                        .iter()
253                        .enumerate()
254                        .skip(1)
255                        .find_map(|(i, claim)| {
256                            if claim.part_idx != 0 {
257                                Some(i - 1)
258                            } else {
259                                None
260                            }
261                        })
262                        .unwrap_or(claims.len() - 1);
263                    last_main_idx_per_proof.push(last_main_idx);
264                    claims.to_device_on(device_ctx).unwrap()
265                })
266                .collect_vec();
267            let lambda_pows = preflights_gpu
268                .iter()
269                .enumerate()
270                .map(|(proof_idx, preflight)| {
271                    preflight
272                        .cpu
273                        .stacking
274                        .lambda
275                        .square()
276                        .powers()
277                        .take(claims[proof_idx].len())
278                        .collect_vec()
279                        .to_device_on(device_ctx)
280                        .unwrap()
281                })
282                .collect_vec();
283
284            let height = if let Some(height) = required_height {
285                if height < num_valid_rows {
286                    return None;
287                }
288                height
289            } else {
290                num_valid_rows.next_power_of_two()
291            };
292            let width = OpeningClaimsCols::<usize>::width();
293            let d_trace = DeviceMatrix::with_capacity_on(height, width, device_ctx);
294            let d_keys_buffer = DeviceBuffer::<F>::with_capacity_on(height, device_ctx);
295
296            let d_claims = claims.iter().map(|buf| buf.as_ptr()).collect_vec();
297            let d_slice_data = blob.slice_data.iter().map(|buf| buf.as_ptr()).collect_vec();
298            let d_precomps = blob.precomps.iter().map(|buf| buf.as_ptr()).collect_vec();
299            let d_lambda_pows = lambda_pows.iter().map(|buf| buf.as_ptr()).collect_vec();
300            let d_records = preflights_gpu
301                .iter()
302                .zip(last_main_idx_per_proof)
303                .map(|(preflight, last_main_idx)| OpeningRecordsPerProof {
304                    tidx_before_column_openings: preflight
305                        .cpu
306                        .batch_constraint
307                        .tidx_before_column_openings
308                        as u32,
309                    last_main_idx: last_main_idx as u32,
310                    lambda: preflight.cpu.stacking.lambda,
311                })
312                .collect_vec()
313                .to_device_on(device_ctx)
314                .unwrap();
315
316            unsafe {
317                let temp_bytes = opening_claims_tracegen_temp_bytes(
318                    d_trace.buffer(),
319                    height,
320                    &d_keys_buffer,
321                    device_ctx.stream.as_raw(),
322                )
323                .unwrap();
324                let d_temp_buffer = DeviceBuffer::<u8>::with_capacity_on(temp_bytes, device_ctx);
325                opening_claims_tracegen(
326                    d_trace.buffer(),
327                    height,
328                    width,
329                    &row_bounds,
330                    d_claims,
331                    d_slice_data,
332                    d_precomps,
333                    d_lambda_pows,
334                    &d_records,
335                    proofs_gpu.len() as u32,
336                    child_vk.system_params.l_skip as u32,
337                    &d_keys_buffer,
338                    &d_temp_buffer,
339                    temp_bytes,
340                    device_ctx.stream.as_raw(),
341                )
342                .unwrap();
343            }
344
345            Some(AirProvingContext::simple_no_pis(d_trace))
346        }
347    }
348}