Skip to main content

sp1_gpu_basefold/
fri.rs

1use itertools::Itertools;
2use std::{marker::PhantomData, sync::Arc};
3
4use slop_algebra::{AbstractExtensionField, AbstractField, ExtensionField, TwoAdicField};
5use slop_alloc::{Buffer, HasBackend};
6use slop_basefold::{BasefoldProof, FriConfig, BATCH_GRINDING_BITS};
7use slop_basefold_prover::{host_fold_even_odd, BasefoldProverError};
8use slop_challenger::{CanObserve, CanSampleBits, FieldChallenger, IopCtx};
9use slop_commit::{Message, Rounds};
10use slop_merkle_tree::MerkleTreeOpeningAndProof;
11use slop_multilinear::{partial_lagrange_blocking, Mle, MultilinearPcsChallenger, Point};
12use slop_tensor::Tensor;
13use sp1_primitives::{SP1ExtensionField, SP1Field};
14
15use sp1_gpu_cudart::{
16    args,
17    sys::{
18        basefold::{
19            batch_koala_bear_base_ext_kernel, batch_koala_bear_base_ext_kernel_flattened,
20            flatten_to_base_koala_bear_base_ext_kernel,
21            transpose_even_odd_koala_bear_base_ext_kernel,
22        },
23        runtime::KernelPtr,
24    },
25    DeviceBuffer, DeviceMle, DeviceTensor, TaskScope,
26};
27use sp1_gpu_merkle_tree::{CudaTcsProver, MerkleTreeProverData, SingleLayerMerkleTreeProverError};
28use sp1_gpu_utils::{Ext, Felt, JaggedTraceMle, TraceDenseData};
29
30use crate::{
31    encode_batch, CudaStackedPcsProverData, DeviceGrindingChallenger, GrindingPowCudaProver,
32    SpparkDftKoalaBear,
33};
34
35/// # Safety
36///
37pub unsafe trait MleBatchKernel<F: TwoAdicField, EF: ExtensionField<F>> {
38    fn batch_mle_kernel() -> KernelPtr;
39}
40
41/// # Safety
42///
43pub unsafe trait RsCodeWordBatchKernel<F: TwoAdicField, EF: ExtensionField<F>> {
44    fn batch_rs_codeword_kernel() -> KernelPtr;
45}
46
47/// # Safety
48pub unsafe trait RsCodeWordTransposeKernel<F: TwoAdicField, EF: ExtensionField<F>> {
49    fn transpose_even_odd_kernel() -> KernelPtr;
50}
51
52/// # Safety
53pub unsafe trait MleFlattenKernel<F: TwoAdicField, EF: ExtensionField<F>> {
54    fn flatten_to_base_kernel() -> KernelPtr;
55}
56
57pub struct FriCudaProver<GC, P, F> {
58    pub tcs_prover: P,
59    pub config: FriConfig<F>,
60    pub log_height: u32,
61    _marker: PhantomData<GC>,
62}
63
64impl<GC: IopCtx<F = Felt, EF = Ext>, P> FriCudaProver<GC, P, GC::F>
65where
66    GC::F: TwoAdicField,
67    GC::EF: ExtensionField<GC::F> + TwoAdicField,
68    P: CudaTcsProver<GC>,
69
70    TaskScope: MleBatchKernel<GC::F, GC::EF>
71        + RsCodeWordBatchKernel<GC::F, GC::EF>
72        + RsCodeWordTransposeKernel<GC::F, GC::EF>
73        + MleFlattenKernel<GC::F, GC::EF>,
74{
75    pub fn new(tcs_prover: P, config: FriConfig<GC::F>, log_height: u32) -> Self {
76        Self { tcs_prover, config, log_height, _marker: PhantomData }
77    }
78
79    pub fn encode_and_commit(
80        &self,
81        use_preprocessed: bool,
82        drop_traces: bool,
83        jagged_trace_mle: &JaggedTraceMle<Felt, TaskScope>,
84    ) -> Result<
85        (<GC as IopCtx>::Digest, CudaStackedPcsProverData<GC>),
86        SingleLayerMerkleTreeProverError,
87    > {
88        let encoder = SpparkDftKoalaBear::default();
89        let scope = jagged_trace_mle.dense().dense.backend().clone();
90
91        let virtual_tensor = if use_preprocessed {
92            jagged_trace_mle.preprocessed_virtual_tensor(self.log_height)
93        } else {
94            jagged_trace_mle.main_virtual_tensor(self.log_height)
95        };
96        let sizes =
97            [virtual_tensor.sizes()[0], 1 << (self.log_height as usize + self.config.log_blowup())];
98
99        let mut dst = Tensor::<GC::F, TaskScope>::with_sizes_in(sizes, scope);
100        unsafe {
101            dst.assume_init();
102        }
103
104        encode_batch(encoder, self.config.log_blowup as u32, virtual_tensor, &mut dst).unwrap();
105
106        // Commit to the tensors.
107
108        let (commitment, tcs_data) = self.tcs_prover.commit_tensors(&dst)?;
109
110        let codeword_mle = if drop_traces { None } else { Some(Arc::new(dst)) };
111        let prover_data = CudaStackedPcsProverData { merkle_tree_tcs_data: tcs_data, codeword_mle };
112
113        Ok((commitment, prover_data))
114    }
115
116    #[allow(clippy::type_complexity)]
117    pub fn batch(
118        &self,
119        batching_coefficients: &Tensor<GC::EF>,
120        mles: &TraceDenseData<GC::F, TaskScope>,
121        codewords: Message<Tensor<Felt, TaskScope>>,
122        evaluation_claims: Vec<GC::EF>,
123    ) -> (Mle<GC::EF, TaskScope>, Tensor<GC::F, TaskScope>, GC::EF) {
124        let log_stacking_height = self.log_height;
125        // Compute all the batch challenge powers.
126        let total_num_polynomials = codewords.iter().map(|c| c.sizes()[0]).sum::<usize>();
127
128        // Compute the random linear combination of the MLEs of the columns of the matrices
129        let num_variables = log_stacking_height;
130        let codeword_size = (codewords.first().unwrap()).sizes()[1];
131        let scope: TaskScope = mles.backend().clone();
132        // All three buffers below are fully overwritten by the kernels, so skip the zero-init.
133        let mut batch_mle = Mle::new(Tensor::<GC::EF, TaskScope>::with_sizes_in(
134            [1, 1 << num_variables],
135            scope.clone(),
136        ));
137        let mut batch_codeword = Tensor::<GC::F, TaskScope>::with_sizes_in(
138            [<GC::EF as AbstractExtensionField<GC::F>>::D, codeword_size],
139            scope.clone(),
140        );
141
142        unsafe {
143            let block_dim = 256;
144            let grid_dim = (1usize << num_variables).div_ceil(block_dim);
145            let batch_size = total_num_polynomials;
146            let powers_device = DeviceBuffer::from_host(batching_coefficients.as_buffer(), &scope)
147                .unwrap()
148                .into_inner();
149            let mle_args = args!(
150                mles.dense.as_ptr(),
151                batch_mle.guts_mut().as_mut_ptr(),
152                powers_device.as_ptr(),
153                (1 << num_variables) as usize,
154                batch_size
155            );
156            batch_mle.assume_init();
157            scope
158                .launch_kernel(TaskScope::batch_mle_kernel(), grid_dim, block_dim, &mle_args, 0)
159                .unwrap();
160        }
161
162        let block_dim = 256;
163        let mut batch_mle_flattened = Mle::new(Tensor::<GC::F, TaskScope>::with_sizes_in(
164            [<GC::EF as AbstractExtensionField<GC::F>>::D, 1 << num_variables],
165            scope.clone(),
166        ));
167        let grid_dim = (1usize << num_variables).div_ceil(block_dim);
168        unsafe {
169            let args = args!(
170                batch_mle.guts().as_ptr(),
171                batch_mle_flattened.guts_mut().as_mut_ptr(),
172                1usize << num_variables
173            );
174            batch_mle_flattened.assume_init();
175            batch_codeword.assume_init();
176            scope
177                .launch_kernel(TaskScope::flatten_to_base_kernel(), grid_dim, block_dim, &args, 0)
178                .unwrap();
179        }
180        let encoder = SpparkDftKoalaBear::default();
181        encode_batch(
182            encoder,
183            self.config.log_blowup as u32,
184            batch_mle_flattened.guts().as_view(),
185            &mut batch_codeword,
186        )
187        .unwrap();
188
189        // Compute the batched evaluation claim.
190        let batch_eval_claim = evaluation_claims
191            .into_iter()
192            .zip(batching_coefficients.as_slice())
193            .map(|(eval, coeff)| eval * *coeff)
194            .sum::<GC::EF>();
195
196        (batch_mle, batch_codeword, batch_eval_claim)
197    }
198
199    #[allow(clippy::type_complexity)]
200    fn commit_phase_round(
201        &self,
202        current_mle: Mle<GC::EF, TaskScope>,
203        current_codeword: Tensor<GC::F, TaskScope>,
204        challenger: &mut GC::Challenger,
205    ) -> Result<
206        (
207            GC::EF,
208            Mle<GC::EF, TaskScope>,
209            Tensor<GC::F, TaskScope>,
210            GC::Digest,
211            Tensor<GC::F, TaskScope>,
212            MerkleTreeProverData<GC::Digest>,
213        ),
214        SingleLayerMerkleTreeProverError,
215    > {
216        // Perform a single round of the FRI commit phase, returning the commitment, folded
217        // codeword, and folding parameter.
218        // On CPU, the current codeword is in row-major form, which means that in order to put
219        // even and odd entries together all we need to do is rehsape it to multiply the number of
220        // columns by 2 and divide the number of rows by 2.
221        let codeword_size = current_codeword.sizes()[1];
222        let batch_size = current_codeword.sizes()[0];
223        let scope = current_codeword.backend().clone();
224
225        let mut leaves = Tensor::with_sizes_in([batch_size * 2, codeword_size / 2], scope.clone());
226        let output_codeword_size = codeword_size / 2;
227        let block_dim = 256;
228        let grid_dim = output_codeword_size.div_ceil(block_dim);
229        unsafe {
230            let args = args!(current_codeword.as_ptr(), leaves.as_mut_ptr(), output_codeword_size);
231            leaves.assume_init();
232            scope
233                .launch_kernel(
234                    TaskScope::transpose_even_odd_kernel(),
235                    grid_dim,
236                    block_dim,
237                    &args,
238                    0,
239                )
240                .unwrap();
241        }
242
243        let (commit, prover_data) = self.tcs_prover.commit_tensors(&leaves)?;
244        // Observe the commitment.
245        challenger.observe(commit);
246
247        let beta: GC::EF = challenger.sample_ext_element();
248
249        // Fold the mle.
250        let folded_mle: Mle<_, TaskScope> = {
251            let device_mle = DeviceMle::from(current_mle);
252            device_mle.fold(beta).into()
253        };
254        let folded_num_variables = folded_mle.num_variables();
255
256        if folded_num_variables < 4 {
257            let current_codeword_transposed =
258                DeviceTensor::from_raw(current_codeword.clone()).transpose();
259            let current_codeword_vec = current_codeword_transposed.to_host().unwrap();
260            let current_codeword_vec =
261                current_codeword_vec.into_buffer().into_extension::<GC::EF>().into_vec();
262            let folded_codeword_vec = host_fold_even_odd(current_codeword_vec, beta);
263            let folded_codeword_storage =
264                Buffer::from(folded_codeword_vec).flatten_to_base::<GC::F>();
265            let mut new_size = current_codeword.sizes().to_vec();
266            new_size[1] /= 2;
267            let folded_codeword =
268                DeviceBuffer::from_host(&folded_codeword_storage, folded_mle.backend())
269                    .unwrap()
270                    .into_inner();
271            let folded_codeword = Tensor::from(folded_codeword).reshape([new_size[1], new_size[0]]);
272            let folded_codeword = DeviceTensor::from_raw(folded_codeword).transpose().into_inner();
273            return Ok((beta, folded_mle, folded_codeword, commit, leaves, prover_data));
274        }
275
276        let folded_height = 1 << folded_num_variables;
277        let mut folded_mle_flattened = Tensor::<GC::F, TaskScope>::with_sizes_in(
278            [<GC::EF as AbstractExtensionField<GC::F>>::D, folded_height],
279            scope.clone(),
280        );
281
282        // Fully overwritten by `encode_batch`.
283        let mut folded_codeword = Tensor::<GC::F, TaskScope>::with_sizes_in(
284            [<GC::EF as AbstractExtensionField<GC::F>>::D, folded_height << self.config.log_blowup],
285            scope.clone(),
286        );
287
288        let block_dim = 256;
289        let grid_dim = folded_height.div_ceil(block_dim);
290        unsafe {
291            let args =
292                args!(folded_mle.guts().as_ptr(), folded_mle_flattened.as_mut_ptr(), folded_height);
293            folded_mle_flattened.assume_init();
294            scope
295                .launch_kernel(TaskScope::flatten_to_base_kernel(), grid_dim, block_dim, &args, 0)
296                .unwrap();
297        }
298        let encoder = SpparkDftKoalaBear::default();
299        encode_batch(
300            encoder,
301            self.config.log_blowup as u32,
302            folded_mle_flattened.as_view(),
303            &mut folded_codeword,
304        )
305        .unwrap();
306
307        Ok((beta, folded_mle, folded_codeword, commit, leaves, prover_data))
308    }
309
310    fn final_poly(&self, final_codeword: Tensor<GC::F, TaskScope>) -> GC::EF {
311        let final_codeword_host = DeviceTensor::from_raw(final_codeword).to_host().unwrap();
312        let final_codeword_transposed = final_codeword_host.transpose();
313        GC::EF::from_base_slice(
314            &final_codeword_transposed.storage.as_slice()
315                [0..(<GC::EF as AbstractExtensionField<GC::F>>::D)],
316        )
317    }
318
319    #[inline]
320    pub fn prove_trusted_evaluations_basefold(
321        &self,
322        mut eval_point: Point<GC::EF>,
323        evaluation_claims: Vec<GC::EF>,
324        mles: &JaggedTraceMle<GC::F, TaskScope>,
325        prover_data: Rounds<&CudaStackedPcsProverData<GC>>,
326        challenger: &mut GC::Challenger,
327    ) -> Result<BasefoldProof<GC>, BasefoldProverError<SingleLayerMerkleTreeProverError>>
328    where
329        GC::Challenger: DeviceGrindingChallenger<Witness = GC::F>,
330    {
331        let scope = mles.dense().dense.backend().clone();
332        let mut codewords: Vec<Arc<Tensor<Felt, TaskScope>>> = Vec::new();
333        for data in prover_data.iter() {
334            if let Some(ref codeword) = data.codeword_mle {
335                codewords.push(codeword.clone());
336            } else {
337                // Codeword was dropped - this is always a main trace.
338                let mut dst = Tensor::<Felt, TaskScope>::with_sizes_in(
339                    [
340                        mles.dense().main_size() >> self.log_height,
341                        1 << (self.log_height as usize + self.config.log_blowup()),
342                    ],
343                    scope.clone(),
344                );
345                unsafe {
346                    dst.assume_init();
347                }
348
349                let encoder = SpparkDftKoalaBear::default();
350                encode_batch(
351                    encoder,
352                    self.config.log_blowup as u32,
353                    mles.main_virtual_tensor(self.log_height),
354                    &mut dst,
355                )
356                .unwrap();
357
358                codewords.push(Arc::new(dst));
359            }
360        }
361
362        let total_num_polynomials = codewords.iter().map(|c| c.sizes()[0]).sum::<usize>();
363        let num_batching_variables = total_num_polynomials.next_power_of_two().ilog2();
364
365        let encoded_messages: Message<_> = codewords.iter().cloned().collect();
366
367        // Grind for batch randomness.
368        let batch_grinding_witness =
369            GrindingPowCudaProver::grind(challenger, BATCH_GRINDING_BITS, &scope);
370
371        let batching_point = challenger.sample_point::<GC::EF>(num_batching_variables);
372        let batching_coefficients = partial_lagrange_blocking(&batching_point);
373
374        // Batch the mles and codewords.
375        let (mle_batch, codeword_batch, batched_eval_claim) =
376            self.batch(&batching_coefficients, mles.dense(), encoded_messages, evaluation_claims);
377        // From this point on, run the BaseFold protocol on the random linear combination codeword,
378        // the random linear combination multilinear, and the random linear combination of the
379        // evaluation claims.
380        let mut current_mle = mle_batch;
381        let mut current_codeword = codeword_batch;
382        // Initialize the vecs that go into a BaseFoldProof.
383        let log_len = current_mle.num_variables();
384        let mut univariate_messages: Vec<[GC::EF; 2]> = vec![];
385        let mut fri_commitments = vec![];
386        let mut commit_phase_data = vec![];
387        let mut current_batched_eval_claim = batched_eval_claim;
388        let mut commit_phase_values = vec![];
389
390        assert_eq!(
391            current_mle.num_variables(),
392            eval_point.dimension() as u32,
393            "eval point dimension mismatch"
394        );
395
396        challenger.observe(Felt::from_canonical_usize(eval_point.dimension()));
397        for _ in 0..eval_point.dimension() {
398            // Compute claims for `g(X_0, X_1, ..., X_{d-1}, 0)` and `g(X_0, X_1, ..., X_{d-1}, 1)`.
399            let last_coord = eval_point.remove_last_coordinate();
400            let zero_values = {
401                use sp1_gpu_cudart::DeviceMle;
402                let device_mle = DeviceMle::from(current_mle.clone());
403                let evals = device_mle.fixed_at_zero(&eval_point);
404                evals.to_host_vec().unwrap()
405            };
406            let zero_val = zero_values[0];
407            let one_val = (current_batched_eval_claim - zero_val) / last_coord + zero_val;
408            let uni_poly = [zero_val, one_val];
409            univariate_messages.push(uni_poly);
410
411            uni_poly.iter().for_each(|elem| challenger.observe_ext_element(*elem));
412
413            // Perform a single round of the FRI commit phase, returning the commitment, folded
414            // codeword, and folding parameter.
415            let (beta, folded_mle, folded_codeword, commitment, leaves, prover_data) = self
416                .commit_phase_round(current_mle, current_codeword, challenger)
417                .map_err(BasefoldProverError::CommitPhaseError)?;
418
419            fri_commitments.push(commitment);
420            commit_phase_data.push(prover_data);
421            commit_phase_values.push(leaves);
422
423            current_mle = folded_mle;
424            current_codeword = folded_codeword;
425            current_batched_eval_claim = zero_val + beta * one_val;
426        }
427
428        let final_poly = self.final_poly(current_codeword);
429        challenger.observe_ext_element(final_poly);
430
431        let fri_config = self.config;
432        let pow_bits = fri_config.proof_of_work_bits;
433        let pow_witness = GrindingPowCudaProver::grind(challenger, pow_bits, &scope);
434        // FRI Query Phase.
435        let query_indices: Vec<usize> = (0..fri_config.num_queries)
436            .map(|_| challenger.sample_bits(log_len as usize + fri_config.log_blowup()))
437            .collect();
438
439        // Open the original polynomials at the query indices.
440        let mut component_polynomials_query_openings_and_proofs = vec![];
441        for (data, codeword) in prover_data.iter().zip(codewords.iter()) {
442            let values = self.tcs_prover.compute_openings_at_indices(codeword, &query_indices);
443            let proof = self
444                .tcs_prover
445                .prove_openings_at_indices(&data.merkle_tree_tcs_data, codeword, &query_indices)
446                .map_err(BasefoldProverError::TcsCommitError)?;
447            let opening = MerkleTreeOpeningAndProof::<GC> { values, proof };
448            component_polynomials_query_openings_and_proofs.push(opening);
449        }
450
451        // Provide openings for the FRI query phase.
452        let mut query_phase_openings_and_proofs = vec![];
453        let mut indices = query_indices;
454        for (leaves, data) in commit_phase_values.into_iter().zip_eq(commit_phase_data) {
455            for index in indices.iter_mut() {
456                *index >>= 1;
457            }
458            let values = self.tcs_prover.compute_openings_at_indices(&leaves, &indices);
459
460            let proof = self
461                .tcs_prover
462                .prove_openings_at_indices(&data, &leaves, &indices)
463                .map_err(BasefoldProverError::TcsCommitError)?;
464            let opening = MerkleTreeOpeningAndProof { values, proof };
465            query_phase_openings_and_proofs.push(opening);
466        }
467
468        Ok(BasefoldProof {
469            univariate_messages,
470            fri_commitments,
471            component_polynomials_query_openings_and_proofs,
472            query_phase_openings_and_proofs,
473            final_poly,
474            pow_witness,
475            batch_grinding_witness,
476        })
477    }
478}
479
480unsafe impl MleBatchKernel<SP1Field, SP1ExtensionField> for TaskScope {
481    fn batch_mle_kernel() -> KernelPtr {
482        unsafe { batch_koala_bear_base_ext_kernel() }
483    }
484}
485
486unsafe impl RsCodeWordBatchKernel<SP1Field, SP1ExtensionField> for TaskScope {
487    fn batch_rs_codeword_kernel() -> KernelPtr {
488        unsafe { batch_koala_bear_base_ext_kernel_flattened() }
489    }
490}
491
492unsafe impl RsCodeWordTransposeKernel<SP1Field, SP1ExtensionField> for TaskScope {
493    fn transpose_even_odd_kernel() -> KernelPtr {
494        unsafe { transpose_even_odd_koala_bear_base_ext_kernel() }
495    }
496}
497
498unsafe impl MleFlattenKernel<SP1Field, SP1ExtensionField> for TaskScope {
499    fn flatten_to_base_kernel() -> KernelPtr {
500        unsafe { flatten_to_base_koala_bear_base_ext_kernel() }
501    }
502}
503
504#[cfg(test)]
505mod tests {
506    use std::sync::Arc;
507
508    use slop_alloc::{CpuBackend, ToHost};
509    use slop_basefold::BasefoldVerifier;
510    use slop_basefold_prover::BasefoldProver;
511    use slop_commit::Message;
512    use slop_futures::queue::WorkerQueue;
513    use slop_merkle_tree::Poseidon2KoalaBear16Prover;
514    use slop_multilinear::{Evaluations, Mle, MleEval};
515    use slop_stacked::interleave_multilinears_with_fixed_rate;
516    use sp1_gpu_cudart::{run_sync_in_place, PinnedBuffer};
517    use sp1_gpu_merkle_tree::{CudaTcsProver, Poseidon2SP1Field16CudaProver};
518    use sp1_gpu_tracegen::CudaTraceGenerator;
519    use sp1_hypercube::prover::{ProverSemaphore, TraceGenerator};
520
521    use sp1_core_machine::io::SP1Stdin;
522    use sp1_gpu_jagged_tracegen::test_utils::tracegen_setup::{
523        self, CORE_MAX_LOG_ROW_COUNT, LOG_STACKING_HEIGHT,
524    };
525    use sp1_gpu_jagged_tracegen::{full_tracegen, CORE_MAX_TRACE_SIZE};
526    use sp1_gpu_utils::{Ext, Felt, TestGC};
527    use sp1_primitives::fri_params::core_fri_config;
528    use sp1_primitives::SP1GlobalContext;
529
530    use super::*;
531
532    #[test]
533    fn test_basefold() {
534        let rt = tokio::runtime::Runtime::new().unwrap();
535        let (machine, record, program) =
536            rt.block_on(tracegen_setup::setup(&test_artifacts::FIBONACCI_ELF, SP1Stdin::new()));
537
538        run_sync_in_place(|scope| {
539            let verifier = BasefoldVerifier::<SP1GlobalContext>::new(core_fri_config(), 2);
540            let old_prover =
541                BasefoldProver::<SP1GlobalContext, Poseidon2KoalaBear16Prover>::new(&verifier);
542
543            let new_cuda_prover = FriCudaProver::<TestGC, _, Felt>::new(
544                Poseidon2SP1Field16CudaProver::new(&scope),
545                verifier.fri_config,
546                LOG_STACKING_HEIGHT,
547            );
548
549            // Generate traces using the host tracegen.
550            let semaphore = ProverSemaphore::new(1);
551            let trace_generator = CudaTraceGenerator::new_in(machine.clone(), scope.clone());
552            let old_traces = rt.block_on(trace_generator.generate_traces(
553                program.clone(),
554                record.clone(),
555                CORE_MAX_LOG_ROW_COUNT as usize,
556                semaphore.clone(),
557            ));
558
559            let preprocessed_traces = old_traces.preprocessed_traces.clone();
560
561            let message = preprocessed_traces
562                .into_iter()
563                .filter_map(|mle| mle.1.into_inner())
564                .map(|x| Clone::clone(x.as_ref()))
565                .collect::<Message<Mle<_, _>>>();
566
567            let host_message: Message<_> = message
568                .clone()
569                .into_iter()
570                .map(|mle| {
571                    let mle = Arc::unwrap_or_clone(mle);
572                    let guts = mle.into_guts();
573                    let device_mle = sp1_gpu_cudart::DeviceMle::from(guts);
574                    device_mle.to_host().unwrap()
575                })
576                .collect();
577
578            let interleaved_message =
579                interleave_multilinears_with_fixed_rate(32, host_message, LOG_STACKING_HEIGHT);
580
581            let interleaved_message =
582                interleaved_message.into_iter().map(|x| x.as_ref().clone()).collect::<Message<_>>();
583
584            let (old_preprocessed_commitment, old_preprocessed_prover_data) =
585                old_prover.commit_mles(interleaved_message.clone()).unwrap();
586
587            let new_semaphore = ProverSemaphore::new(1);
588            let capacity = CORE_MAX_TRACE_SIZE as usize;
589            let buffer = PinnedBuffer::<Felt>::with_capacity(capacity);
590            let queue = Arc::new(WorkerQueue::new(vec![buffer]));
591            let buffer = rt.block_on(queue.pop()).unwrap();
592            let (_, new_traces, _, _) = rt.block_on(full_tracegen(
593                &machine,
594                program,
595                Arc::new(record),
596                &buffer,
597                CORE_MAX_TRACE_SIZE as usize,
598                LOG_STACKING_HEIGHT,
599                CORE_MAX_LOG_ROW_COUNT,
600                &scope,
601                new_semaphore,
602                false,
603            ));
604
605            let (new_preprocessed_commit, new_preprocessed_prover_data) =
606                new_cuda_prover.encode_and_commit(true, false, &new_traces).unwrap();
607
608            assert_eq!(new_preprocessed_commit, old_preprocessed_commitment);
609
610            let (new_main_commit, new_main_prover_data) =
611                new_cuda_prover.encode_and_commit(false, false, &new_traces).unwrap();
612            let message = old_traces
613                .main_trace_data
614                .traces
615                .into_iter()
616                .filter_map(|mle| mle.1.into_inner())
617                .map(|x| Clone::clone(x.as_ref()))
618                .collect::<Message<Mle<_, _>>>();
619
620            let mut host_message = Vec::new();
621            for mle in message.into_iter() {
622                let mle = Arc::unwrap_or_clone(mle);
623                let guts = mle.into_guts();
624                let device_mle = sp1_gpu_cudart::DeviceMle::from(guts);
625                let mle_host = device_mle.to_host().unwrap();
626                host_message.push(mle_host);
627            }
628
629            let host_message = host_message.into_iter().collect::<Message<Mle<Felt, CpuBackend>>>();
630
631            let interleaved_message_2 =
632                interleave_multilinears_with_fixed_rate(32, host_message, LOG_STACKING_HEIGHT);
633
634            let (old_main_commitment, old_main_prover_data) =
635                old_prover.commit_mles(interleaved_message_2.clone()).unwrap();
636
637            assert_eq!(new_main_commit, old_main_commitment);
638
639            let mut rng = rand::thread_rng();
640
641            let eval_point_host = Point::<Ext>::rand(&mut rng, LOG_STACKING_HEIGHT);
642
643            let evaluation_claims_1: Vec<_> = interleaved_message
644                .clone()
645                .into_iter()
646                .map(|mle| mle.eval_at(&eval_point_host))
647                .collect();
648
649            let evaluation_claims_1 = Evaluations { round_evaluations: evaluation_claims_1 };
650
651            let evaluation_claims_2: Vec<_> = interleaved_message_2
652                .clone()
653                .into_iter()
654                .map(|mle| mle.eval_at(&eval_point_host))
655                .collect();
656
657            let host_evaluation_claims_1: Vec<MleEval<Ext, CpuBackend>> = evaluation_claims_1
658                .round_evaluations
659                .iter()
660                .map(|mle| mle.to_host().unwrap())
661                .collect();
662
663            let host_evaluation_claims_2: Vec<MleEval<Ext, CpuBackend>> =
664                evaluation_claims_2.iter().map(|mle| mle.to_host().unwrap()).collect();
665
666            let flattened_evaluation_claims = vec![
667                MleEval::new(
668                    host_evaluation_claims_1
669                        .into_iter()
670                        .flat_map(|x: MleEval<Ext, CpuBackend>| x.evaluations().storage.to_vec())
671                        .collect(),
672                ),
673                MleEval::new(
674                    host_evaluation_claims_2
675                        .into_iter()
676                        .flat_map(|x: MleEval<Ext, CpuBackend>| x.evaluations().storage.to_vec())
677                        .collect(),
678                ),
679            ];
680
681            let evaluation_claims_2 = Evaluations { round_evaluations: evaluation_claims_2 };
682
683            let mut challenger = SP1GlobalContext::default_challenger();
684
685            scope.synchronize_blocking().unwrap();
686            let now = std::time::Instant::now();
687
688            let basefold_proof = old_prover
689                .prove_trusted_mle_evaluations(
690                    eval_point_host.clone(),
691                    vec![interleaved_message, interleaved_message_2].into_iter().collect(),
692                    vec![evaluation_claims_1.clone(), evaluation_claims_2.clone()]
693                        .into_iter()
694                        .collect(),
695                    vec![old_preprocessed_prover_data, old_main_prover_data].into_iter().collect(),
696                    &mut challenger,
697                )
698                .unwrap();
699
700            scope.synchronize_blocking().unwrap();
701            tracing::info!("Old proof time: {:?}", now.elapsed());
702
703            let mut challenger = SP1GlobalContext::default_challenger();
704
705            let flat_evaluation_claims: Vec<Ext> = evaluation_claims_1
706                .round_evaluations
707                .iter()
708                .chain(evaluation_claims_2.round_evaluations.iter())
709                .flat_map(|mle_eval| mle_eval.iter().copied())
710                .collect();
711
712            scope.synchronize_blocking().unwrap();
713
714            let now = std::time::Instant::now();
715
716            let new_basefold_proof = new_cuda_prover
717                .prove_trusted_evaluations_basefold(
718                    eval_point_host.clone(),
719                    flat_evaluation_claims,
720                    &new_traces,
721                    [&new_preprocessed_prover_data, &new_main_prover_data].into_iter().collect(),
722                    &mut challenger,
723                )
724                .unwrap();
725
726            scope.synchronize_blocking().unwrap();
727            tracing::info!("New proof time: {:?}", now.elapsed());
728
729            // Because the batch grinding is non-deterministic between CPU and GPU, the
730            // grinding witnesses may differ, causing all subsequent proof values (batching
731            // point, univariate messages, etc.) to diverge. Instead of comparing proof
732            // components directly, we verify both proofs independently.
733
734            verifier
735                .verify_mle_evaluations(
736                    &[old_preprocessed_commitment, old_main_commitment],
737                    eval_point_host.clone(),
738                    &flattened_evaluation_claims,
739                    &basefold_proof,
740                    &mut SP1GlobalContext::default_challenger(),
741                )
742                .unwrap();
743
744            verifier
745                .verify_mle_evaluations(
746                    &[new_preprocessed_commit, new_main_commit],
747                    eval_point_host,
748                    &flattened_evaluation_claims,
749                    &new_basefold_proof,
750                    &mut SP1GlobalContext::default_challenger(),
751                )
752                .unwrap();
753        })
754        .unwrap();
755    }
756}