1use std::ffi::c_void;
7
8use openvm_cuda_common::{
9 copy::cuda_memcpy_on,
10 d_buffer::DeviceBuffer,
11 error::{CudaError, MemCopyError},
12 stream::GpuDeviceCtx,
13};
14use openvm_stark_backend::{
15 p3_challenger::{CanObserve, CanSample},
16 FiatShamirTranscript, StarkProtocolConfig,
17};
18use openvm_stark_sdk::config::baby_bear_poseidon2::poseidon2_perm;
19use p3_baby_bear::default_babybear_poseidon2_16;
20use p3_field::{PrimeCharacteristicRing, PrimeField32};
21use p3_symmetric::Permutation;
22
23use crate::types::{Challenger, Digest, CHUNK, F, SC, WIDTH};
24
25pub(crate) fn validate_gpu_grind_bits(bits: usize) -> Result<(), GrindError> {
26 if bits >= u32::BITS as usize || (1u64 << bits) >= u64::from(F::ORDER_U32) {
27 return Err(CudaError::new(1).into());
28 }
29 Ok(())
30}
31
32#[repr(C)]
40#[derive(Clone, Debug)]
41pub struct DeviceSpongeState {
42 pub state: [F; WIDTH],
44 pub absorb_idx: u32,
46 pub sample_idx: u32,
48}
49
50impl Default for DeviceSpongeState {
51 fn default() -> Self {
52 Self {
53 state: [F::default(); WIDTH],
54 absorb_idx: 0,
55 sample_idx: 0,
56 }
57 }
58}
59
60impl DeviceSpongeState {
61 #[inline]
65 pub fn observe(&mut self, value: F) {
66 self.state[self.absorb_idx as usize] = value;
67 self.absorb_idx += 1;
68 if self.absorb_idx == CHUNK as u32 {
69 poseidon2_perm().permute_mut(&mut self.state);
70 self.absorb_idx = 0;
71 self.sample_idx = CHUNK as u32;
72 }
73 }
74
75 #[inline]
79 pub fn sample(&mut self) -> F {
80 if self.absorb_idx != 0 || self.sample_idx == 0 {
81 poseidon2_perm().permute_mut(&mut self.state);
82 self.absorb_idx = 0;
83 self.sample_idx = CHUNK as u32;
84 }
85 self.sample_idx -= 1;
86 self.state[self.sample_idx as usize]
87 }
88}
89
90impl FiatShamirTranscript<SC> for DeviceSpongeState {
91 #[inline]
92 fn observe(&mut self, value: F) {
93 DeviceSpongeState::observe(self, value);
94 }
95
96 #[inline]
97 fn sample(&mut self) -> F {
98 DeviceSpongeState::sample(self)
99 }
100
101 #[inline]
102 fn observe_commit(&mut self, digest: Digest) {
103 for x in digest {
104 self.observe(x);
105 }
106 }
107}
108
109#[derive(Debug)]
142pub struct DuplexSpongeGpu {
143 host: Challenger,
145 device: DeviceBuffer<DeviceSpongeState>,
147}
148
149impl Default for DuplexSpongeGpu {
150 fn default() -> Self {
151 Self::new()
152 }
153}
154
155impl Clone for DuplexSpongeGpu {
156 fn clone(&self) -> Self {
157 Self {
160 host: self.host.clone(),
161 device: DeviceBuffer::new(),
162 }
163 }
164}
165
166impl DuplexSpongeGpu {
167 pub fn new() -> Self {
169 Self {
170 host: Challenger::new(default_babybear_poseidon2_16()),
171 device: DeviceBuffer::new(),
172 }
173 }
174
175 pub fn is_device_allocated(&self) -> bool {
177 !self.device.is_empty()
178 }
179
180 fn ensure_device_allocated(&mut self, device_ctx: &GpuDeviceCtx) {
182 if self.device.is_empty() {
183 self.device = DeviceBuffer::with_capacity_on(1, device_ctx);
184 }
185 }
186
187 pub fn sync_h2d(&mut self, device_ctx: &GpuDeviceCtx) -> Result<(), MemCopyError> {
194 self.ensure_device_allocated(device_ctx);
195
196 let mut device_state = DeviceSpongeState {
202 state: self.host.sponge_state,
203 absorb_idx: self.host.input_buffer.len() as u32,
204 sample_idx: self.host.output_buffer.len() as u32,
205 };
206
207 for (i, &val) in self.host.input_buffer.iter().enumerate() {
211 device_state.state[i] = val;
212 }
213
214 unsafe {
218 cuda_memcpy_on::<false, true>(
219 self.device.as_mut_ptr() as *mut c_void,
220 &device_state as *const DeviceSpongeState as *const c_void,
221 std::mem::size_of::<DeviceSpongeState>(),
222 device_ctx,
223 )
224 }
225 }
226
227 pub fn device_ptr(&self) -> Option<*const DeviceSpongeState> {
232 if self.device.is_empty() {
233 None
234 } else {
235 Some(self.device.as_ptr())
236 }
237 }
238
239 pub fn device_ptr_mut(&mut self) -> Option<*mut DeviceSpongeState> {
244 if self.device.is_empty() {
245 None
246 } else {
247 Some(self.device.as_mut_ptr())
248 }
249 }
250
251 pub fn grind_gpu(&mut self, bits: usize, device_ctx: &GpuDeviceCtx) -> Result<F, GrindError> {
268 validate_gpu_grind_bits(bits)?;
269 if bits == 0 {
271 return Ok(F::ZERO);
272 }
273 self.sync_h2d(device_ctx)?;
275
276 let witness_u32 = unsafe {
278 crate::cuda::sponge::sponge_grind(
279 self.device.as_ptr(),
280 bits as u32,
281 F::ORDER_U32 - 1,
282 device_ctx,
283 )?
284 };
285
286 let witness = F::from_u32(witness_u32);
287
288 debug_assert!(self.clone().check_witness(bits, witness));
291 self.host.observe(witness);
292 let _: F = self.host.sample(); Ok(witness)
295 }
296}
297
298#[derive(Debug, thiserror::Error)]
300pub enum GrindError {
301 #[error("Memory copy error: {0}")]
302 MemCopy(#[from] MemCopyError),
303
304 #[error("CUDA error: {0}")]
305 Cuda(#[from] CudaError),
306
307 #[error("Failed to find PoW witness within search space")]
308 WitnessNotFound,
309}
310
311impl FiatShamirTranscript<SC> for DuplexSpongeGpu {
312 #[inline]
313 fn observe(&mut self, value: F) {
314 self.host.observe(value);
315 }
316
317 #[inline]
318 fn sample(&mut self) -> F {
319 self.host.sample()
320 }
321
322 #[inline]
323 fn observe_commit(&mut self, digest: Digest) {
324 for x in digest {
325 self.observe(x);
326 }
327 }
328}
329
330pub trait GpuFiatShamirTranscript<Config: StarkProtocolConfig>:
336 FiatShamirTranscript<Config>
337{
338 fn grind_gpu(
345 &mut self,
346 bits: usize,
347 device_ctx: &GpuDeviceCtx,
348 ) -> Result<Config::F, GrindError>;
349}
350
351impl GpuFiatShamirTranscript<SC> for DuplexSpongeGpu {
352 fn grind_gpu(&mut self, bits: usize, device_ctx: &GpuDeviceCtx) -> Result<F, GrindError> {
353 DuplexSpongeGpu::grind_gpu(self, bits, device_ctx)
354 }
355}
356
357#[cfg(test)]
358mod tests {
359 use std::time::Instant;
360
361 use openvm_cuda_common::{
362 common::get_device,
363 stream::{CudaStream, GpuDeviceCtx, StreamGuard},
364 };
365 use openvm_stark_sdk::config::baby_bear_poseidon2::default_duplex_sponge;
366 use p3_field::PrimeCharacteristicRing;
367
368 use super::*;
369 use crate::prelude::SC;
370
371 fn test_ctx() -> GpuDeviceCtx {
372 GpuDeviceCtx {
373 device_id: get_device().unwrap() as u32,
374 stream: StreamGuard::new(CudaStream::new_non_blocking().unwrap()),
375 }
376 }
377
378 #[test]
379 fn test_device_sponge_state_size() {
380 let expected_size = std::mem::size_of::<[F; WIDTH]>() + std::mem::size_of::<u32>() + std::mem::size_of::<u32>(); assert_eq!(
386 std::mem::size_of::<DeviceSpongeState>(),
387 expected_size,
388 "DeviceSpongeState size mismatch - check repr(C) and padding"
389 );
390 }
391
392 #[test]
393 fn test_device_sponge_state_alignment() {
394 assert!(
396 std::mem::align_of::<DeviceSpongeState>() >= 4,
397 "DeviceSpongeState should be at least 4-byte aligned"
398 );
399 }
400
401 #[test]
402 fn test_default_state() {
403 let state = DeviceSpongeState::default();
404 assert_eq!(state.absorb_idx, 0);
405 assert_eq!(state.sample_idx, 0);
406 for elem in state.state.iter() {
407 assert_eq!(*elem, F::default());
408 }
409 }
410
411 #[test]
412 fn test_sponge_gpu_new() {
413 let sponge = DuplexSpongeGpu::new();
414 assert!(!sponge.is_device_allocated());
415 }
416
417 #[test]
418 fn test_device_sponge_state_matches_duplex_sponge() {
419 let mut device_state = DeviceSpongeState::default();
421 let mut duplex_sponge = default_duplex_sponge();
422
423 for i in 0..20 {
425 let val = F::from_u32(i * 42 + 17);
426 device_state.observe(val);
427 FiatShamirTranscript::<SC>::observe(&mut duplex_sponge, val);
428 }
429
430 for _ in 0..10 {
431 let device_sample = device_state.sample();
432 let duplex_sample = FiatShamirTranscript::<SC>::sample(&mut duplex_sponge);
433 assert_eq!(device_sample, duplex_sample);
434 }
435
436 for i in 0..5 {
438 let val = F::from_u32(i * 100);
439 device_state.observe(val);
440 FiatShamirTranscript::<SC>::observe(&mut duplex_sponge, val);
441
442 let device_sample = device_state.sample();
443 let duplex_sample = FiatShamirTranscript::<SC>::sample(&mut duplex_sponge);
444 assert_eq!(device_sample, duplex_sample);
445 }
446
447 for _ in 0..15 {
449 let device_sample = device_state.sample();
450 let duplex_sample = FiatShamirTranscript::<SC>::sample(&mut duplex_sponge);
451 assert_eq!(device_sample, duplex_sample);
452 }
453 }
454
455 #[test]
456 fn test_sponge_gpu_uses_host_transcript() {
457 let mut gpu_sponge = DuplexSpongeGpu::default();
458 let mut cpu_sponge = default_duplex_sponge();
459
460 for i in 0..10 {
462 let val = F::from_u32(i * 42 + 17);
463 gpu_sponge.observe(val);
464 FiatShamirTranscript::<SC>::observe(&mut cpu_sponge, val);
465 }
466
467 for _ in 0..5 {
468 let gpu_sample = gpu_sponge.sample();
469 let cpu_sample = FiatShamirTranscript::<SC>::sample(&mut cpu_sponge);
470 assert_eq!(gpu_sample, cpu_sample);
471 }
472 }
473
474 #[test]
482 fn test_grind_cpu_vs_gpu() {
483 let device_ctx = test_ctx();
484 {
486 let mut warmup = DuplexSpongeGpu::default();
487 let _ = warmup.grind_gpu(8, &device_ctx);
488 }
489
490 let bit_counts = [8, 12, 16, 18, 20]
492 .iter()
493 .flat_map(|x| std::iter::repeat_n(*x, 5))
494 .collect::<Vec<_>>();
495
496 eprintln!("\n{}", "=".repeat(60));
497 eprintln!("Grinding Performance: CPU vs GPU");
498 eprintln!("{}", "=".repeat(60));
499 eprintln!(
500 "{:>6} {:>12} {:>12} {:>10}",
501 "bits", "CPU (ms)", "GPU (ms)", "speedup"
502 );
503 eprintln!("{:->6} {:->12} {:->12} {:->10}", "", "", "", "");
504
505 let mut seed = 265;
506 for bits in bit_counts {
507 let mut cpu_sponge = default_duplex_sponge();
508 let mut gpu_sponge = DuplexSpongeGpu::default();
509
510 for _ in 0..5 {
512 let val = F::from_u32(seed);
513 seed += 228;
514 FiatShamirTranscript::<SC>::observe(&mut cpu_sponge, val);
515 gpu_sponge.observe(val);
516 }
517
518 let cpu_start = Instant::now();
520 let cpu_witness = FiatShamirTranscript::<SC>::grind(&mut cpu_sponge, bits);
521 let cpu_time = cpu_start.elapsed();
522
523 let gpu_start = Instant::now();
525 let gpu_witness = gpu_sponge
526 .grind_gpu(bits, &device_ctx)
527 .expect("GPU grinding failed");
528 let gpu_time = gpu_start.elapsed();
529
530 let speedup = cpu_time.as_secs_f64() / gpu_time.as_secs_f64();
534
535 eprintln!(
536 "{:>6} {:>12.2} {:>12.2} {:>10.2}x",
537 bits,
538 cpu_time.as_secs_f64() * 1000.0,
539 gpu_time.as_secs_f64() * 1000.0,
540 speedup
541 );
542
543 let _ = (cpu_witness, gpu_witness); }
547
548 eprintln!("{}\n", "=".repeat(60));
549 }
550}