openvm_cuda_common/
common.rs1use std::{
2 ffi::c_void,
3 sync::atomic::{AtomicU64, Ordering},
4};
5
6use crate::error::{check, CudaError};
7
8#[link(name = "cudart")]
9extern "C" {
10 fn cudaFree(dev_ptr: *mut c_void) -> i32;
11 fn cudaGetDevice(device: *mut i32) -> i32;
12 fn cudaSetDevice(device: i32) -> i32;
13 fn cudaDeviceReset() -> i32;
14}
15
16static DEVICE_RESET_EPOCH: AtomicU64 = AtomicU64::new(0);
17
18pub fn get_device() -> Result<i32, CudaError> {
19 let mut device = 0;
20 unsafe {
21 check(cudaGetDevice(&mut device))?;
22 }
23 assert!(device >= 0);
24 Ok(device)
25}
26
27pub fn device_reset_epoch() -> u64 {
28 DEVICE_RESET_EPOCH.load(Ordering::Acquire)
29}
30
31pub fn set_device_by_id(device: i32) -> Result<(), CudaError> {
32 assert!(device >= 0);
33 unsafe {
34 check(cudaSetDevice(device))?;
35 check(cudaFree(std::ptr::null_mut()))?;
37 }
38 Ok(())
39}
40
41pub fn set_device() -> Result<i32, CudaError> {
42 let device = get_device()?;
43 set_device_by_id(device)?;
44 Ok(device)
45}
46
47pub fn reset_device() -> Result<(), CudaError> {
48 check(unsafe { cudaDeviceReset() })?;
49 DEVICE_RESET_EPOCH.fetch_add(1, Ordering::AcqRel);
50 Ok(())
51}