openvm_recursion_circuit/stacking/claims/
trace.rs1use std::borrow::BorrowMut;
2
3use itertools::{izip, Itertools};
4use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, EF, F};
5use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
6use p3_matrix::dense::RowMajorMatrix;
7
8use crate::{
9 stacking::{
10 claims::air::StackingClaimsCols,
11 utils::{compute_coefficients, get_stacked_slice_data},
12 },
13 tracegen::{RowMajorChip, StandardTracegenCtx},
14};
15
16pub struct StackingClaimsTraceGenerator;
17
18impl RowMajorChip<F> for StackingClaimsTraceGenerator {
19 type Ctx<'a> = StandardTracegenCtx<'a>;
20
21 #[tracing::instrument(level = "trace", skip_all)]
22 fn generate_trace(
23 &self,
24 ctx: &Self::Ctx<'_>,
25 required_height: Option<usize>,
26 ) -> Option<RowMajorMatrix<F>> {
27 let vk = ctx.vk;
28 let proofs = ctx.proofs;
29 let preflights = ctx.preflights;
30 debug_assert_eq!(proofs.len(), preflights.len());
31
32 let w_stack = vk.inner.params.w_stack;
33 let width = StackingClaimsCols::<usize>::width();
34 let minimum_height = proofs.len() * w_stack;
36 let height = if let Some(height) = required_height {
37 if height < minimum_height {
38 return None;
39 }
40 height
41 } else {
42 minimum_height.next_power_of_two()
43 };
44
45 let mut trace = vec![F::ZERO; height * width];
46 let mut chunks = trace.chunks_mut(width);
47
48 for (proof_idx, (proof, preflight)) in proofs.iter().zip(preflights).enumerate() {
49 let claims = proof
50 .stacking_proof
51 .stacking_openings
52 .iter()
53 .enumerate()
54 .flat_map(|(commit_idx, openings)| {
55 openings
56 .iter()
57 .enumerate()
58 .map(move |(stacked_col_idx, opening)| {
59 (commit_idx, stacked_col_idx, opening)
60 })
61 })
62 .collect_vec();
63 let stacked_slices =
64 get_stacked_slice_data(vk, &preflight.proof_shape.sorted_trace_vdata);
65
66 let coeffs = compute_coefficients(
67 proof,
68 &stacked_slices,
69 &preflight.stacking.sumcheck_rnd,
70 &preflight.batch_constraint.sumcheck_rnd,
71 &preflight.stacking.lambda,
72 vk.inner.params.l_skip,
73 vk.inner.params.n_stack,
74 )
75 .0
76 .into_iter()
77 .flatten()
78 .collect_vec();
79
80 let num_valid = claims.len();
81 debug_assert!(
82 num_valid <= w_stack,
83 "proof {proof_idx} has {num_valid} stacking claims but w_stack = {w_stack}"
84 );
85 let proof_idx_value = F::from_usize(proof_idx);
86
87 let initial_tidx = preflight.stacking.intermediate_tidx[2];
88
89 let mu = preflight.stacking.stacking_batching_challenge;
90 let mu_pows = mu.powers().take(num_valid).collect_vec();
91
92 let mu_pow_witness = preflight.stacking.mu_pow_witness;
94 let mu_pow_sample = preflight.stacking.mu_pow_sample;
95
96 let mut final_s_eval = EF::ZERO;
97 let mut whir_claim = EF::ZERO;
98
99 for (idx, (&(commit_idx, stacked_col_idx, &claim), coeff)) in
100 izip!(&claims, coeffs).enumerate()
101 {
102 let chunk = chunks.next().unwrap();
103 let cols: &mut StackingClaimsCols<F> = chunk.borrow_mut();
104
105 cols.proof_idx = proof_idx_value;
106 cols.is_valid = F::ONE;
107 cols.is_first = F::from_bool(idx == 0);
108 cols.is_last = F::from_bool(num_valid == w_stack && idx + 1 == num_valid);
109
110 cols.commit_idx = F::from_usize(commit_idx);
111 cols.stacked_col_idx = F::from_usize(stacked_col_idx);
112 cols.global_col_idx = F::from_usize(idx);
113
114 cols.tidx = F::from_usize(initial_tidx + (D_EF * idx));
115 cols.mu.copy_from_slice(mu.as_basis_coefficients_slice());
116 cols.mu_pow
117 .copy_from_slice(mu_pows[idx].as_basis_coefficients_slice());
118
119 cols.stacking_claim
120 .copy_from_slice(claim.as_basis_coefficients_slice());
121
122 cols.mu_pow_witness = mu_pow_witness;
124 cols.mu_pow_sample = mu_pow_sample;
125
126 cols.claim_coefficient
127 .copy_from_slice(coeff.as_basis_coefficients_slice());
128 final_s_eval += claim * coeff;
129 cols.final_s_eval
130 .copy_from_slice(final_s_eval.as_basis_coefficients_slice());
131
132 whir_claim += mu_pows[idx] * claim;
133 cols.whir_claim
134 .copy_from_slice(whir_claim.as_basis_coefficients_slice());
135 }
136
137 for idx in num_valid..w_stack {
139 let chunk = chunks.next().unwrap();
140 let cols: &mut StackingClaimsCols<F> = chunk.borrow_mut();
141
142 cols.proof_idx = proof_idx_value;
143 cols.is_padding = F::ONE;
144 cols.is_last = F::from_bool(idx + 1 == w_stack);
145 cols.global_col_idx = F::from_usize(idx);
146 }
147 }
148
149 let padding_proof_idx = F::from_usize(proofs.len());
150 let mut chunks = chunks.peekable();
151
152 while let Some(chunk) = chunks.next() {
153 let cols: &mut StackingClaimsCols<F> = chunk.borrow_mut();
154 cols.proof_idx = padding_proof_idx;
155 if chunks.peek().is_none() {
156 cols.is_last = F::ONE;
157 }
158 }
159
160 Some(RowMajorMatrix::new(trace, width))
161 }
162}
163
164#[cfg(feature = "cuda")]
165pub(crate) mod cuda {
166 use openvm_cuda_backend::{base::DeviceMatrix, GpuBackend};
167 use openvm_cuda_common::{copy::MemCopyH2D, d_buffer::DeviceBuffer};
168 use openvm_stark_backend::prover::AirProvingContext;
169
170 use super::*;
171 use crate::{
172 stacking::{
173 cuda_abi::{
174 stacking_claims_tracegen, stacking_claims_tracegen_temp_bytes,
175 ClaimsRecordsPerProof, StackingClaim,
176 },
177 cuda_tracegen::StackingBlob,
178 },
179 tracegen::{cuda::StandardTracegenGpuCtx, ModuleChip},
180 };
181
182 pub struct StackingClaimsTraceGeneratorGpu;
183
184 impl ModuleChip<GpuBackend> for StackingClaimsTraceGeneratorGpu {
185 type Ctx<'a> = (StandardTracegenGpuCtx<'a>, &'a StackingBlob);
186
187 fn generate_proving_ctx(
188 &self,
189 ctx: &Self::Ctx<'_>,
190 required_height: Option<usize>,
191 ) -> Option<openvm_stark_backend::prover::AirProvingContext<GpuBackend>> {
192 let proofs_gpu = ctx.0.proofs;
193 let preflights_gpu = ctx.0.preflights;
194 let device_ctx = ctx.0.device_ctx;
195 let blob = ctx.1;
196 let w_stack = ctx.0.vk.system_params.w_stack;
197
198 let mut row_bounds = Vec::with_capacity(proofs_gpu.len());
199 let claims = proofs_gpu
200 .iter()
201 .enumerate()
202 .map(|(proof_idx, proof)| {
203 let claims = proof
204 .cpu
205 .stacking_proof
206 .stacking_openings
207 .iter()
208 .enumerate()
209 .flat_map(|(commit_idx, openings)| {
210 openings
211 .iter()
212 .enumerate()
213 .map(move |(stacked_col_idx, opening)| StackingClaim {
214 commit_idx: commit_idx as u32,
215 stacked_col_idx: stacked_col_idx as u32,
216 claim: *opening,
217 })
218 })
219 .collect_vec();
220
221 let num_valid = claims.len();
222 assert!(
223 num_valid <= w_stack,
224 "proof {proof_idx} has {num_valid} stacking claims but w_stack = {w_stack}"
225 );
226 row_bounds.push(((proof_idx + 1) * w_stack) as u32);
227
228 claims.to_device_on(device_ctx).unwrap()
229 })
230 .collect_vec();
231
232 let mu_pows = preflights_gpu
233 .iter()
234 .enumerate()
235 .map(|(proof_idx, preflight)| {
236 let mu = preflight.cpu.stacking.stacking_batching_challenge;
237 mu.powers()
238 .take(claims[proof_idx].len())
239 .collect_vec()
240 .to_device_on(device_ctx)
241 .unwrap()
242 })
243 .collect_vec();
244
245 let minimum_height = proofs_gpu.len() * w_stack;
246 let height = if let Some(height) = required_height {
247 if height < minimum_height {
248 return None;
249 }
250 height
251 } else {
252 minimum_height.next_power_of_two()
253 };
254 let width = StackingClaimsCols::<usize>::width();
255 let d_trace = DeviceMatrix::with_capacity_on(height, width, device_ctx);
256
257 let d_claims = claims.iter().map(|buf| buf.as_ptr()).collect_vec();
258 let d_coeffs = blob.coeffs.iter().map(|buf| buf.as_ptr()).collect_vec();
259 let d_mu_pows = mu_pows.iter().map(|buf| buf.as_ptr()).collect_vec();
260 let d_records = preflights_gpu
261 .iter()
262 .enumerate()
263 .map(|(proof_idx, preflight)| ClaimsRecordsPerProof {
264 initial_tidx: preflight.cpu.stacking.intermediate_tidx[2] as u32,
265 num_valid: claims[proof_idx].len() as u32,
266 mu: preflight.cpu.stacking.stacking_batching_challenge,
267 mu_pow_witness: preflight.cpu.stacking.mu_pow_witness,
268 mu_pow_sample: preflight.cpu.stacking.mu_pow_sample,
269 })
270 .collect_vec()
271 .to_device_on(device_ctx)
272 .unwrap();
273
274 unsafe {
275 let temp_bytes = stacking_claims_tracegen_temp_bytes(
276 d_trace.buffer(),
277 height,
278 device_ctx.stream.as_raw(),
279 )
280 .unwrap();
281 let d_temp_buffer = DeviceBuffer::<u8>::with_capacity_on(temp_bytes, device_ctx);
282 stacking_claims_tracegen(
283 d_trace.buffer(),
284 height,
285 width,
286 &row_bounds,
287 d_claims,
288 d_coeffs,
289 d_mu_pows,
290 &d_records,
291 proofs_gpu.len() as u32,
292 &d_temp_buffer,
293 temp_bytes,
294 device_ctx.stream.as_raw(),
295 )
296 .unwrap();
297 }
298
299 Some(AirProvingContext::simple_no_pis(d_trace))
300 }
301 }
302}