openvm_recursion_circuit/stacking/opening/
trace.rs1use 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}