Skip to main content

openvm_cuda_common/
common.rs

1use 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        // Force primary-context initialization on the selected device.
36        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}