mod context;
mod loader;
mod registry;
#[cfg(target_family = "windows")]
mod veh;
pub use context::{create_temp_context, destroy_temp_context};
pub use registry::CudaAllocationRegistry;
use core::ffi::c_void;
use core::sync::atomic::Ordering;
use registry::{register_cuda_ptr_in, unregister_cuda_ptr_in};
trait CudaAllocOps {
fn registry() -> &'static CudaAllocationRegistry;
fn alloc_sym() -> *mut c_void;
fn free_sym() -> *mut c_void;
unsafe fn raw_alloc(alloc_sym: *mut c_void, size: usize) -> *mut u8;
unsafe fn raw_free(free_sym: *mut c_void, ptr: *mut u8) -> core::ffi::c_int;
#[inline]
unsafe fn post_register(_ptr: *mut u8, _size: usize) {}
}
#[inline]
unsafe fn cuda_allocate<Ops: CudaAllocOps>(size: usize) -> *mut u8 {
unsafe { loader::init_cuda() };
let alloc_sym = Ops::alloc_sym();
let free_sym = Ops::free_sym();
if alloc_sym.is_null() || free_sym.is_null() {
return core::ptr::null_mut();
}
let ptr = unsafe { Ops::raw_alloc(alloc_sym, size) };
if ptr.is_null() {
return core::ptr::null_mut();
}
if register_cuda_ptr_in(Ops::registry(), ptr) {
unsafe { Ops::post_register(ptr, size) };
return ptr;
}
let _rollback_status = unsafe { Ops::raw_free(free_sym, ptr) };
core::ptr::null_mut()
}
#[inline]
unsafe fn cuda_deallocate<Ops: CudaAllocOps>(ptr: *mut u8) -> bool {
if !unregister_cuda_ptr_in(Ops::registry(), ptr) {
return false;
}
let free_sym = Ops::free_sym();
if free_sym.is_null() {
return false;
}
unsafe { Ops::raw_free(free_sym, ptr) == 0 }
}
unsafe fn managed_raw_alloc(alloc_sym: *mut c_void, size: usize) -> *mut u8 {
type CuMemAllocManagedFn = unsafe extern "system" fn(*mut u64, usize, u32) -> core::ffi::c_int;
let cu_mem_alloc_managed: CuMemAllocManagedFn = unsafe { core::mem::transmute(alloc_sym) };
let mut dptr: u64 = 0;
let res = unsafe { cu_mem_alloc_managed(&mut dptr, size, 0x01) };
if res == 0 && dptr != 0 {
dptr as *mut u8
} else {
core::ptr::null_mut()
}
}
unsafe fn managed_raw_free(free_sym: *mut c_void, ptr: *mut u8) -> core::ffi::c_int {
type CuMemFreeFn = unsafe extern "system" fn(u64) -> core::ffi::c_int;
let cu_mem_free: CuMemFreeFn = unsafe { core::mem::transmute(free_sym) };
unsafe { cu_mem_free(ptr as u64) }
}
macro_rules! impl_device_tier_backend {
($backend:ty) => {
impl mnemosyne_core::MemoryBackend for $backend {
const SUPPORTS_PAGE_RESET: bool = CudaDeviceBackend::SUPPORTS_PAGE_RESET;
const SUPPORTS_MAKE_GUARD: bool = CudaDeviceBackend::SUPPORTS_MAKE_GUARD;
const SUPPORTS_DECOMMIT: bool = CudaDeviceBackend::SUPPORTS_DECOMMIT;
const ENABLE_CPU_CACHE: bool = CudaDeviceBackend::ENABLE_CPU_CACHE;
#[inline(always)]
unsafe fn allocate(size: usize) -> *mut u8 {
unsafe { allocate_device_tier(size) }
}
#[inline(always)]
unsafe fn deallocate(ptr: *mut u8, size: usize) -> bool {
unsafe { deallocate_device_tier(ptr, size) }
}
}
};
}
macro_rules! impl_cuda_memory_backend {
($backend:ty) => {
impl mnemosyne_core::MemoryBackend for $backend {
#[inline]
unsafe fn allocate(size: usize) -> *mut u8 {
unsafe { super::cuda_allocate::<Self>(size) }
}
#[inline]
unsafe fn deallocate(ptr: *mut u8, _size: usize) -> bool {
unsafe { super::cuda_deallocate::<Self>(ptr) }
}
}
};
}
mod device;
mod pinned;
mod unified;
pub use device::{CudaDeviceBackend, CudaGddrBackend, CudaHbmBackend};
pub use pinned::CudaHostPinnedBackend;
pub use unified::CudaUnifiedBackend;
pub fn is_cuda_available() -> bool {
unsafe { loader::init_cuda() };
!loader::CU_MEM_ALLOC_MANAGED
.load(Ordering::Acquire)
.is_null()
}