use std::fmt;
use std::hash::{Hash, Hasher};
use std::io::Write;
use std::sync::{Arc, Mutex, OnceLock};
use cubecl::client::ComputeClient;
use cubecl::stream_id::StreamId;
use cubecl::Runtime;
use cubecl_cuda::{CudaDevice, CudaRuntime as CubeclCudaRuntime};
use cubecl_runtime::config::{CubeClRuntimeConfig, RuntimeConfig};
use cudarc::cublas::sys as cublas_sys;
use cudarc::driver::result::DriverError;
use cudarc::driver::sys::{CUcontext, CUdevice, CUresult};
use cudarc::runtime::{result as cuda_result, sys as cuda_sys, sys::cudaStream_t};
use tenferro_tensor::AllocationDomainId;
use super::device::{
cuda_devices, unavailable_device_error, CudaDeviceError, CudaDeviceId, CudaDeviceInfo,
};
use super::identity::GpuExtensionCapability;
pub fn gpu_available() -> bool {
let library_present = std::panic::catch_unwind(|| {
unsafe { cudarc::driver::sys::is_culib_present() }
})
.unwrap_or(false);
if !library_present {
return false;
}
let Ok(devices) = cuda_devices() else {
return false;
};
let Some(device_id) = devices.first().map(|device| device.id()) else {
return false;
};
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let Ok(runtime) = CudaRuntime::new(device_id) else {
return false;
};
runtime.synchronize().is_ok()
}))
.unwrap_or(false)
}
pub(crate) struct RawContextRestore {
saved_device: Result<i32, cudarc::runtime::result::RuntimeError>,
saved_context: Result<Option<CUcontext>, cudarc::driver::result::DriverError>,
op: &'static str,
}
impl RawContextRestore {
pub(crate) fn enter(op: &'static str, device: i32, context: CUcontext) -> crate::Result<Self> {
let saved_device = cudarc::runtime::result::device::get();
let saved_context = cudarc::driver::result::ctx::get_current();
cudarc::runtime::result::device::set(device)
.map_err(|err| crate::Error::backend_source(op, err))?;
if let Err(err) = unsafe { cudarc::driver::result::ctx::set_current(context) } {
if let Ok(previous_device) = saved_device {
let _ = cudarc::runtime::result::device::set(previous_device);
}
match saved_context {
Ok(Some(previous)) => {
let _ = unsafe { cudarc::driver::result::ctx::set_current(previous) };
}
Ok(None) => {
let _ =
unsafe { cudarc::driver::result::ctx::set_current(std::ptr::null_mut()) };
}
Err(_) => {}
}
return Err(crate::Error::backend_source(op, err));
}
Ok(Self {
saved_device,
saved_context,
op,
})
}
fn restore(&self) {
let mut stderr = std::io::stderr();
if let Ok(device) = self.saved_device {
if let Err(err) = cudarc::runtime::result::device::set(device) {
let _ = writeln!(
stderr,
"tenferro-gpu: failed to restore CUDA device during {}: {err:?}",
self.op
);
}
}
match self.saved_context {
Ok(Some(context)) => {
if let Err(err) = unsafe { cudarc::driver::result::ctx::set_current(context) } {
let _ = writeln!(
stderr,
"tenferro-gpu: failed to restore CUDA context during {}: {err:?}",
self.op
);
}
}
Ok(None) => {
if let Err(err) =
unsafe { cudarc::driver::result::ctx::set_current(std::ptr::null_mut()) }
{
let _ = writeln!(
stderr,
"tenferro-gpu: failed to clear CUDA context during {}: {err:?}",
self.op
);
}
}
Err(_) => {}
}
}
}
impl Drop for RawContextRestore {
fn drop(&mut self) {
self.restore();
}
}
#[derive(Clone, Debug)]
pub struct CudaRuntimeIdentity {
marker: Arc<u8>,
}
impl CudaRuntimeIdentity {
fn fresh() -> Self {
Self {
marker: Arc::new(0),
}
}
}
impl PartialEq for CudaRuntimeIdentity {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.marker, &other.marker)
}
}
impl Eq for CudaRuntimeIdentity {}
impl Hash for CudaRuntimeIdentity {
fn hash<H: Hasher>(&self, state: &mut H) {
state.write_usize(Arc::as_ptr(&self.marker) as usize);
}
}
#[derive(Clone)]
pub struct CudaRuntime {
inner: Arc<CudaRuntimeState>,
}
struct CudaRuntimeState {
client: ComputeClient<CubeclCudaRuntime>,
device_id: CudaDeviceId,
device_ordinal: usize,
device_info: CudaDeviceInfo,
primary_context: CudaPrimaryContext,
identity: CudaRuntimeIdentity,
allocation_domain: AllocationDomainId,
raw_streams: Box<[OnceLock<u64>]>,
cublas_handles: Box<[Mutex<Option<CublasStreamHandle>>]>,
pinned_scalar: Mutex<PinnedScalarSlot>,
}
struct CublasStreamHandle(cublas_sys::cublasHandle_t);
struct PinnedScalarSlot {
ptr: *mut std::ffi::c_void,
}
pub(crate) const PINNED_SCALAR_BYTES: usize = 16;
unsafe impl Send for CudaRuntimeState {}
unsafe impl Sync for CudaRuntimeState {}
impl fmt::Debug for CudaRuntime {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CudaRuntime")
.field("device_id", &self.inner.device_id)
.finish_non_exhaustive()
}
}
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_source("cubecl_runtime_init", 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_id: CudaDeviceId) -> Result<Self, CudaDeviceError> {
let device_ordinal = usize::try_from(device_id.ordinal()).map_err(|source| {
cuda_initialization_error(device_id, "convert_device_ordinal", source)
})?;
let cuda_ordinal = i32::try_from(device_id.ordinal()).map_err(|source| {
cuda_initialization_error(device_id, "convert_cuda_ordinal", source)
})?;
cudarc::driver::result::init()
.map_err(|source| cuda_initialization_error(device_id, "initialize_driver", source))?;
let cuda_device = match cudarc::driver::result::device::get(cuda_ordinal) {
Ok(cuda_device) => cuda_device,
Err(source) if is_invalid_device_lookup(source) => {
return Err(unavailable_device_error(device_id, cuda_devices()?));
}
Err(source) => {
return Err(cuda_initialization_error(device_id, "get_device", source));
}
};
let primary_context = CudaPrimaryContext::retain(cuda_device).map_err(|source| {
cuda_initialization_error(device_id, "retain_primary_context", source)
})?;
unsafe { cudarc::driver::result::ctx::set_current(primary_context.context()) }.map_err(
|source| cuda_initialization_error(device_id, "set_current_context", source),
)?;
cudarc::runtime::result::device::set(cuda_ordinal)
.map_err(|source| cuda_initialization_error(device_id, "set_device", source))?;
let device = CudaDevice::new(device_ordinal);
let client = CubeclCudaRuntime::client(&device);
let discovered = cuda_devices()?;
let device_info = discovered
.iter()
.find(|info| info.id() == device_id)
.cloned()
.ok_or_else(|| unavailable_device_error(device_id, discovered))?;
Ok(Self {
inner: Arc::new(CudaRuntimeState {
client,
device_id,
device_ordinal,
device_info,
primary_context,
identity: CudaRuntimeIdentity::fresh(),
allocation_domain: AllocationDomainId::fresh(),
raw_streams: (0..cubecl_stream_slots())
.map(|_| OnceLock::new())
.collect(),
cublas_handles: (0..cubecl_stream_slots())
.map(|_| Mutex::new(None))
.collect(),
pinned_scalar: Mutex::new(PinnedScalarSlot {
ptr: std::ptr::null_mut(),
}),
}),
})
}
pub(crate) fn client(&self) -> &ComputeClient<CubeclCudaRuntime> {
&self.inner.client
}
pub fn device_id(&self) -> CudaDeviceId {
self.inner.device_id
}
pub fn device_info(&self) -> &CudaDeviceInfo {
&self.inner.device_info
}
pub fn allocation_domain(&self) -> AllocationDomainId {
self.inner.allocation_domain
}
pub fn supports_extension(&self, capability: GpuExtensionCapability) -> bool {
capabilities_for_device(capability)
}
pub(crate) fn device_ordinal(&self) -> usize {
self.inner.device_ordinal
}
pub(crate) fn primary_context(&self) -> CUcontext {
self.inner.primary_context.context()
}
pub fn with_current_context<R>(
&self,
op: &'static str,
f: impl FnOnce() -> R,
) -> crate::Result<R> {
let device_ordinal = i32::try_from(self.device_ordinal())
.map_err(|source| crate::Error::backend_source(op, source))?;
let _guard = RawContextRestore::enter(op, device_ordinal, self.primary_context())?;
Ok(f())
}
pub(crate) fn flush_cubecl(&self, op: &'static str) -> crate::Result<()> {
self.client()
.flush()
.map_err(|err| crate::Error::backend_source(op, err))
}
pub fn runtime_identity(&self) -> CudaRuntimeIdentity {
self.inner.identity.clone()
}
pub(crate) fn allocation_domain_id(&self) -> AllocationDomainId {
self.inner.allocation_domain
}
pub(crate) fn set_current_cuda_context(&self, op: &'static str) -> crate::Result<()> {
self.inner.set_current_cuda_context(op)
}
pub(crate) fn raw_cuda_stream(&self) -> crate::Result<u64> {
self.inner.raw_cuda_stream()
}
pub(crate) fn synchronize_raw_stream(
&self,
stream: u64,
op: &'static str,
) -> crate::Result<()> {
self.inner.synchronize_raw_stream(stream, op)
}
pub(crate) fn with_cublas_handle<R>(
&self,
op: &'static str,
pointer_mode: cublas_sys::cublasPointerMode_t,
cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
execute: impl FnOnce(cublas_sys::cublasHandle_t) -> crate::Result<R>,
) -> crate::Result<R> {
self.inner
.with_cublas_handle(op, pointer_mode, cross_stream_handles, execute)
}
pub(crate) fn finish_vendor_enqueue<R>(
&self,
op: &'static str,
cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
result: crate::Result<R>,
) -> crate::Result<R> {
self.inner
.finish_vendor_enqueue(op, cross_stream_handles, result)
}
pub(crate) fn stream_slot(&self) -> usize {
self.inner.stream_slot()
}
pub(crate) fn stream_slot_count(&self) -> usize {
self.inner.raw_streams.len()
}
pub(crate) fn is_current_stream_slot(&self, handle: &cubecl_runtime::server::Handle) -> bool {
self.inner.stream_slot_for(handle.stream) == self.inner.stream_slot()
}
pub(crate) fn download_scalar_bytes(
&self,
device_addr: u64,
out: &mut [u8],
op: &'static str,
retained: cubecl_runtime::server::Handle,
) -> crate::Result<()> {
self.inner
.download_scalar_bytes(device_addr, out, op, retained)
}
pub fn synchronize(&self) -> crate::Result<()> {
self.inner.synchronize()
}
}
impl CudaRuntimeState {
fn stream_slot(&self) -> usize {
self.stream_slot_for(StreamId::current())
}
fn stream_slot_for(&self, stream_id: StreamId) -> usize {
stream_id.value as usize % self.raw_streams.len()
}
fn set_current_cuda_context(&self, op: &'static str) -> crate::Result<()> {
if let Ok(Some(current)) = cudarc::driver::result::ctx::get_current() {
if current == self.primary_context.context() {
return Ok(());
}
}
let device_ordinal = i32::try_from(self.device_id.ordinal())
.map_err(|source| crate::Error::backend_source(op, source))?;
cudarc::runtime::result::device::set(device_ordinal)
.map_err(|err| crate::Error::backend_source(op, err))?;
unsafe { cudarc::driver::result::ctx::set_current(self.primary_context.context()) }
.map_err(|err| crate::Error::backend_source(op, err))
}
fn raw_cuda_stream(&self) -> crate::Result<u64> {
let stream_id = StreamId::current();
let slot = self.stream_slot();
if let Some(&stream) = self.raw_streams[slot].get() {
return Ok(stream);
}
let stream = self
.client
.with_server(move |server| {
server
.raw_stream(stream_id)
.map(|stream| stream as u64)
.map_err(|err| crate::Error::backend_source("raw_cuda_stream", err))
})
.ok_or_else(|| {
crate::Error::runtime_state("raw_cuda_stream", "CubeCL server is unavailable")
})??;
Ok(*self.raw_streams[slot].get_or_init(|| stream))
}
fn synchronize(&self) -> crate::Result<()> {
const OP: &str = "cubecl_runtime_synchronize";
let stream = self.raw_cuda_stream()?;
self.synchronize_raw_stream(stream, OP)
}
fn synchronize_raw_stream(&self, stream: u64, op: &'static str) -> crate::Result<()> {
self.set_current_cuda_context(op)?;
unsafe { cuda_result::stream::synchronize(stream as usize as cudaStream_t) }
.map_err(|err| crate::Error::backend_source(op, err))
}
fn retire_initialized_streams(&self) -> bool {
const OP: &str = "cuda_runtime_drop";
if let Err(error) = self.set_current_cuda_context(OP) {
report_cuda_runtime_drop_error(&error);
return false;
}
let mut retired = true;
for stream in &self.raw_streams {
let Some(&stream) = stream.get() else {
continue;
};
if let Err(source) =
unsafe { cuda_result::stream::synchronize(stream as usize as cudaStream_t) }
{
retired = false;
report_cuda_runtime_drop_error(&crate::Error::backend_source(OP, source));
}
}
retired
}
fn with_cublas_handle<R>(
&self,
op: &'static str,
pointer_mode: cublas_sys::cublasPointerMode_t,
cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
execute: impl FnOnce(cublas_sys::cublasHandle_t) -> crate::Result<R>,
) -> crate::Result<R> {
let poisoned = || crate::Error::runtime_state(op, "cuBLAS handle cache lock poisoned");
let slot = self.stream_slot();
let mut cached = self.cublas_handles[slot].lock().map_err(|_| poisoned())?;
let handle = match *cached {
Some(ref handle) => handle.0,
None => {
if !cublas_library_present() {
return Err(crate::Error::io_source(op, CublasLibraryMissing));
}
let stream = self.raw_cuda_stream()? as usize as cublas_sys::cudaStream_t;
let mut raw = std::ptr::null_mut();
check_cublas(op, "cublasCreate", unsafe {
cublas_sys::cublasCreate_v2(&mut raw)
})?;
if let Err(err) = check_cublas(op, "cublasSetStream", unsafe {
cublas_sys::cublasSetStream_v2(raw, stream)
}) {
let _ = unsafe { cublas_sys::cublasDestroy_v2(raw) };
return Err(err);
}
*cached = Some(CublasStreamHandle(raw));
raw
}
};
check_cublas(op, "cublasSetPointerMode", unsafe {
cublas_sys::cublasSetPointerMode_v2(handle, pointer_mode)
})?;
let result = execute(handle);
self.finish_vendor_enqueue(op, cross_stream_handles, result)
}
fn finish_vendor_enqueue<R>(
&self,
op: &'static str,
cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
result: crate::Result<R>,
) -> crate::Result<R> {
if cross_stream_handles.is_empty() {
return result;
}
let retirement = self.synchronize();
match (result, retirement) {
(Ok(value), Ok(())) => Ok(value),
(Err(error), Ok(())) => Err(error),
(Ok(_), Err(retirement)) => {
std::mem::forget(cross_stream_handles);
Err(crate::Error::backend_source(op, retirement))
}
(Err(error), Err(_retirement)) => {
std::mem::forget(cross_stream_handles);
Err(error)
}
}
}
fn download_scalar_bytes(
&self,
device_addr: u64,
out: &mut [u8],
op: &'static str,
retained: cubecl_runtime::server::Handle,
) -> crate::Result<()> {
if out.len() > PINNED_SCALAR_BYTES {
return Err(crate::Error::Internal(format!(
"pinned scalar staging supports at most {PINNED_SCALAR_BYTES} bytes, got {}",
out.len()
)));
}
self.set_current_cuda_context(op)?;
let stream = self.raw_cuda_stream()? as usize as cudaStream_t;
let mut slot = self
.pinned_scalar
.lock()
.map_err(|_| crate::Error::runtime_state(op, "pinned scalar staging lock poisoned"))?;
if slot.ptr.is_null() {
let mut ptr = std::ptr::null_mut();
unsafe {
cuda_sys::cudaHostAlloc(
&mut ptr,
PINNED_SCALAR_BYTES,
cuda_sys::cudaHostAllocDefault,
)
}
.result()
.map_err(|err| crate::Error::backend_source(op, err))?;
slot.ptr = ptr;
}
let src = super::interop::cuda_device_ptr_from_addr(device_addr, op)?;
let staging = unsafe { std::slice::from_raw_parts_mut(slot.ptr.cast::<u8>(), out.len()) };
let completed = unsafe { cuda_result::memcpy_dtoh_async(staging, src, stream) }
.and_then(|()| unsafe { cuda_result::stream::synchronize(stream) });
if let Err(err) = completed {
std::mem::forget(retained);
slot.ptr = std::ptr::null_mut();
return Err(crate::Error::backend_source(op, err));
}
out.copy_from_slice(staging);
Ok(())
}
fn release_cuda_library_resources(&mut self) {
for cached in &self.cublas_handles {
if let Ok(mut handle) = cached.lock() {
let Some(handle) = handle.take() else {
continue;
};
if let Err(err) = unsafe { cublas_sys::cublasDestroy_v2(handle.0) }.result() {
report_cuda_resource_release_error("cuBLAS handle", &err);
}
}
}
if let Ok(mut slot) = self.pinned_scalar.lock() {
if !slot.ptr.is_null() {
if let Err(err) = unsafe { cuda_sys::cudaFreeHost(slot.ptr) }.result() {
report_cuda_resource_release_error("pinned scalar staging", &err);
}
slot.ptr = std::ptr::null_mut();
}
}
}
}
fn cubecl_stream_slots() -> usize {
usize::from(CubeClRuntimeConfig::get().streaming.max_streams.max(1))
}
#[derive(Debug, thiserror::Error)]
#[error(
"cuBLAS shared library not found; ensure `LD_LIBRARY_PATH` includes the CUDA toolkit library directory"
)]
struct CublasLibraryMissing;
fn cublas_library_present() -> bool {
use std::sync::OnceLock;
static PRESENT: OnceLock<bool> = OnceLock::new();
*PRESENT.get_or_init(|| unsafe { cublas_sys::is_culib_present() })
}
pub(super) fn check_cublas(
op: &'static str,
call: &'static str,
status: cublas_sys::cublasStatus_t,
) -> crate::Result<()> {
if matches!(status, cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS) {
Ok(())
} else {
Err(super::error::provider_status(
op,
"cuBLAS",
call,
status as i32,
))
}
}
#[cold]
fn report_cuda_resource_release_error(what: &'static str, err: &impl fmt::Debug) {
eprintln!("tenferro-gpu: failed to release {what} during Drop: {err:?}");
}
fn is_invalid_device_lookup(source: DriverError) -> bool {
source.0 == CUresult::CUDA_ERROR_INVALID_DEVICE
}
pub(crate) fn capabilities_for_device(capability: GpuExtensionCapability) -> bool {
!matches!(capability, GpuExtensionCapability::PeerCopy)
}
fn cuda_initialization_error<E>(
device: CudaDeviceId,
operation: &'static str,
source: E,
) -> CudaDeviceError
where
E: std::error::Error + Send + Sync + 'static,
{
CudaDeviceError::Initialization {
device,
operation,
source: Box::new(source),
}
}
impl Drop for CudaRuntimeState {
fn drop(&mut self) {
if self.retire_initialized_streams() {
self.release_cuda_library_resources();
}
}
}
#[cfg(test)]
mod tests;