use std::fmt;
use cubecl::client::ComputeClient;
use cubecl::stream_id::StreamId;
use cubecl::Runtime;
use cubecl_cuda::{CudaDevice, CudaRuntime as CubeclCudaRuntime};
use cudarc::driver::sys::{CUcontext, CUdevice};
use cudarc::runtime::{result as cuda_result, sys::cudaStream_t};
pub fn gpu_available() -> bool {
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let device = CudaDevice::new(0);
let _ = CubeclCudaRuntime::client(&device);
}))
.is_ok()
}
pub struct CudaRuntime {
client: ComputeClient<CubeclCudaRuntime>,
device_ordinal: usize,
cuda_device: CUdevice,
cuda_context: CUcontext,
}
impl fmt::Debug for CudaRuntime {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CudaRuntime")
.field("device_ordinal", &self.device_ordinal)
.finish_non_exhaustive()
}
}
unsafe impl Send for CudaRuntime {}
impl CudaRuntime {
pub fn new(device_ordinal: usize) -> crate::Result<Self> {
cudarc::runtime::result::device::set(device_ordinal as i32).map_err(|err| {
crate::Error::backend_failure(
"cubecl_runtime_init",
format!("failed to set CUDA runtime device: {err:?}"),
)
})?;
cudarc::driver::result::init().map_err(|err| {
crate::Error::backend_failure(
"cubecl_runtime_init",
format!("failed to initialize CUDA driver: {err:?}"),
)
})?;
let cuda_device =
cudarc::driver::result::device::get(device_ordinal as i32).map_err(|err| {
crate::Error::backend_failure(
"cubecl_runtime_init",
format!("failed to obtain CUDA device {device_ordinal}: {err:?}"),
)
})?;
let cuda_context = unsafe { cudarc::driver::result::primary_ctx::retain(cuda_device) }
.map_err(|err| {
crate::Error::backend_failure(
"cubecl_runtime_init",
format!("failed to retain CUDA primary context: {err:?}"),
)
})?;
unsafe { cudarc::driver::result::ctx::set_current(cuda_context) }.map_err(|err| {
crate::Error::backend_failure(
"cubecl_runtime_init",
format!("failed to set CUDA primary context current: {err:?}"),
)
})?;
let device = CudaDevice::new(device_ordinal);
let client = CubeclCudaRuntime::client(&device);
Ok(Self {
client,
device_ordinal,
cuda_device,
cuda_context,
})
}
pub(crate) fn client(&self) -> &ComputeClient<CubeclCudaRuntime> {
&self.client
}
pub fn device_ordinal(&self) -> usize {
self.device_ordinal
}
#[doc(hidden)]
pub fn set_current_cuda_context(&self, op: &'static str) -> crate::Result<()> {
cudarc::runtime::result::device::set(self.device_ordinal as i32).map_err(|err| {
crate::Error::backend_failure(op, format!("failed to set CUDA runtime device: {err:?}"))
})?;
unsafe { cudarc::driver::result::ctx::set_current(self.cuda_context) }.map_err(|err| {
crate::Error::backend_failure(
op,
format!("failed to activate CUDA primary context: {err:?}"),
)
})
}
pub(crate) fn raw_cuda_stream(&self) -> crate::Result<u64> {
self.client
.with_server(|server| {
server
.raw_stream(StreamId::current())
.map(|stream| stream as u64)
.map_err(|err| {
crate::Error::backend_failure("raw_cuda_stream", format!("{err:?}"))
})
})
.ok_or_else(|| {
crate::Error::backend_failure("raw_cuda_stream", "with_server returned None")
})?
}
pub fn synchronize(&self) -> crate::Result<()> {
const OP: &str = "cubecl_runtime_synchronize";
self.set_current_cuda_context(OP)?;
let stream = self.raw_cuda_stream()? as usize as cudaStream_t;
unsafe { cuda_result::stream::synchronize(stream) }.map_err(|err| {
crate::Error::backend_failure(OP, format!("CUDA stream synchronize failed: {err:?}"))
})
}
}
impl Drop for CudaRuntime {
fn drop(&mut self) {
let _ = self.synchronize();
let _ = unsafe { cudarc::driver::result::primary_ctx::release(self.cuda_device) };
}
}