openvm_recursion_circuit/stacking/claims/
trace.rs

1use 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        // Each proof gets exactly w_stack rows (valid + padding).
35        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            // μ PoW witness and sample from preflight
93            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                // μ PoW columns (only used on last row, but set for all rows for simplicity)
123                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            // Padding rows (fill up to w_stack)
138            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}