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
35pub unsafe trait MleBatchKernel<F: TwoAdicField, EF: ExtensionField<F>> {
38 fn batch_mle_kernel() -> KernelPtr;
39}
40
41pub unsafe trait RsCodeWordBatchKernel<F: TwoAdicField, EF: ExtensionField<F>> {
44 fn batch_rs_codeword_kernel() -> KernelPtr;
45}
46
47pub unsafe trait RsCodeWordTransposeKernel<F: TwoAdicField, EF: ExtensionField<F>> {
49 fn transpose_even_odd_kernel() -> KernelPtr;
50}
51
52pub 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 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 let total_num_polynomials = codewords.iter().map(|c| c.sizes()[0]).sum::<usize>();
127
128 let num_variables = log_stacking_height;
130 let codeword_size = (codewords.first().unwrap()).sizes()[1];
131 let scope: TaskScope = mles.backend().clone();
132 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 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 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 challenger.observe(commit);
246
247 let beta: GC::EF = challenger.sample_ext_element();
248
249 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 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 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 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 let (mle_batch, codeword_batch, batched_eval_claim) =
376 self.batch(&batching_coefficients, mles.dense(), encoded_messages, evaluation_claims);
377 let mut current_mle = mle_batch;
381 let mut current_codeword = codeword_batch;
382 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 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 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 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 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 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 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 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}