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,
primary_context: CudaPrimaryContext,
}
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 {}
struct CudaPrimaryContext {
cuda_device: CUdevice,
cuda_context: CUcontext,
}
impl CudaPrimaryContext {
fn retain(cuda_device: CUdevice) -> crate::Result<Self> {
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:?}"),
)
})?;
Ok(Self {
cuda_device,
cuda_context,
})
}
fn context(&self) -> CUcontext {
self.cuda_context
}
}
impl Drop for CudaPrimaryContext {
fn drop(&mut self) {
if let Err(err) = unsafe { cudarc::driver::result::primary_ctx::release(self.cuda_device) }
{
report_cuda_primary_context_release_error(&err);
}
}
}
#[cold]
fn report_cuda_primary_context_release_error(err: &impl fmt::Debug) {
eprintln!("tenferro-gpu: failed to release CUDA primary context during Drop: {err:?}");
}
#[cold]
fn report_cuda_runtime_drop_error(err: &crate::Error) {
eprintln!("tenferro-gpu: failed to synchronize CUDA runtime during Drop: {err}");
}
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 primary_context = CudaPrimaryContext::retain(cuda_device)?;
unsafe { cudarc::driver::result::ctx::set_current(primary_context.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,
primary_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.primary_context.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) {
if let Err(err) = self.synchronize() {
report_cuda_runtime_drop_error(&err);
}
}
}