use core::ffi::{CStr, c_void};
use core::sync::atomic::{AtomicPtr, AtomicU8, Ordering};
#[cfg(any(target_family = "windows", target_family = "unix"))]
use core::ffi::c_char;
#[cfg(target_family = "windows")]
unsafe extern "system" {
fn LoadLibraryA(lpLibFileName: *const c_char) -> *mut c_void;
fn GetProcAddress(hModule: *mut c_void, lpProcName: *const c_char) -> *mut c_void;
fn SwitchToThread() -> i32;
}
#[cfg(target_family = "unix")]
unsafe extern "C" {
fn dlopen(filename: *const c_char, flag: core::ffi::c_int) -> *mut c_void;
fn dlsym(handle: *mut c_void, symbol: *const c_char) -> *mut c_void;
fn sched_yield() -> core::ffi::c_int;
}
#[cfg(target_family = "unix")]
const RTLD_LAZY: core::ffi::c_int = 1;
static CUDA_LIB: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
pub(super) static CU_MEM_ALLOC_MANAGED: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
pub(super) static CU_MEM_FREE: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
pub(super) static CU_MEM_HOST_ALLOC: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
pub(super) static CU_MEM_FREE_HOST: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
pub(super) static CU_MEM_ADVISE: AtomicPtr<c_void> = AtomicPtr::new(core::ptr::null_mut());
static CUDA_INIT_STATE: AtomicU8 = AtomicU8::new(CUDA_UNINITIALIZED);
const CUDA_UNINITIALIZED: u8 = 0;
const CUDA_INITIALIZING: u8 = 1;
const CUDA_INITIALIZED: u8 = 2;
const INIT_SPIN_LIMIT: u32 = 1024;
pub(super) unsafe fn cuda_library() -> *mut c_void {
let cached = CUDA_LIB.load(Ordering::Acquire);
if !cached.is_null() {
return cached;
}
let lib: *mut c_void = {
#[cfg(target_family = "windows")]
{
unsafe { LoadLibraryA(c"nvcuda.dll".as_ptr()) }
}
#[cfg(target_family = "unix")]
{
let p = unsafe { dlopen(c"libcuda.so".as_ptr(), RTLD_LAZY) };
if p.is_null() {
unsafe { dlopen(c"libcuda.so.1".as_ptr(), RTLD_LAZY) }
} else {
p
}
}
#[cfg(not(any(target_family = "windows", target_family = "unix")))]
{
core::ptr::null_mut()
}
};
if lib.is_null() {
return core::ptr::null_mut();
}
match CUDA_LIB.compare_exchange(
core::ptr::null_mut(),
lib,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => lib,
Err(existing) => existing,
}
}
pub(super) unsafe fn resolve_sym(lib: *mut c_void, name: &CStr) -> *mut c_void {
#[cfg(target_family = "windows")]
{
unsafe { GetProcAddress(lib, name.as_ptr()) }
}
#[cfg(target_family = "unix")]
{
unsafe { dlsym(lib, name.as_ptr()) }
}
#[cfg(not(any(target_family = "windows", target_family = "unix")))]
{
let _ = (lib, name);
core::ptr::null_mut()
}
}
fn yield_thread() {
#[cfg(target_family = "windows")]
{
let _switched = unsafe { SwitchToThread() };
}
#[cfg(target_family = "unix")]
{
let _yielded = unsafe { sched_yield() };
}
#[cfg(not(any(target_family = "windows", target_family = "unix")))]
{
core::hint::spin_loop();
}
}
pub(super) unsafe fn init_cuda() {
if CUDA_INIT_STATE.load(Ordering::Acquire) == CUDA_INITIALIZED {
return;
}
if CUDA_INIT_STATE
.compare_exchange(
CUDA_UNINITIALIZED,
CUDA_INITIALIZING,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
unsafe { init_cuda_once() };
CUDA_INIT_STATE.store(CUDA_INITIALIZED, Ordering::Release);
return;
}
let mut spins: u32 = 0;
while CUDA_INIT_STATE.load(Ordering::Acquire) != CUDA_INITIALIZED {
if spins < INIT_SPIN_LIMIT {
spins += 1;
core::hint::spin_loop();
} else {
yield_thread();
}
}
}
#[cfg(target_family = "windows")]
use super::veh::run_cu_init_isolated;
#[cfg(not(target_family = "windows"))]
unsafe fn run_cu_init_isolated(init_sym: *mut c_void) -> i32 {
type CuInitFn = unsafe extern "system" fn(u32) -> core::ffi::c_int;
let cu_init: CuInitFn = unsafe { core::mem::transmute(init_sym) };
unsafe { cu_init(0) }
}
unsafe fn init_cuda_once() {
let lib = unsafe { cuda_library() };
if lib.is_null() {
return;
}
let init_sym = unsafe { resolve_sym(lib, c"cuInit") };
let alloc_sym = unsafe { resolve_sym(lib, c"cuMemAllocManaged") };
let free_sym = unsafe { resolve_sym(lib, c"cuMemFree") };
let host_alloc_sym = unsafe { resolve_sym(lib, c"cuMemHostAlloc") };
let free_host_sym = unsafe { resolve_sym(lib, c"cuMemFreeHost") };
let advise_sym = unsafe { resolve_sym(lib, c"cuMemAdvise") };
if init_sym.is_null() || alloc_sym.is_null() || free_sym.is_null() {
return;
}
if unsafe { run_cu_init_isolated(init_sym) } != 0 {
return;
}
CU_MEM_ALLOC_MANAGED.store(alloc_sym, Ordering::Release);
CU_MEM_FREE.store(free_sym, Ordering::Release);
if !host_alloc_sym.is_null() {
CU_MEM_HOST_ALLOC.store(host_alloc_sym, Ordering::Release);
}
if !free_host_sym.is_null() {
CU_MEM_FREE_HOST.store(free_host_sym, Ordering::Release);
}
if !advise_sym.is_null() {
CU_MEM_ADVISE.store(advise_sym, Ordering::Release);
}
}