Skip to main content

mnemosyne_backend/backends/cuda/
pinned.rs

1use core::ffi::c_void;
2use core::sync::atomic::Ordering;
3
4use super::registry::CUDA_HOST_PINNED_ALLOCATIONS;
5use super::{CudaAllocOps, CudaAllocationRegistry, loader};
6
7/// A memory backend allocating CUDA page-locked (pinned) host memory.
8///
9/// `allocate` returns null on failure (driver unavailable, driver allocation
10/// failure, or registry full); there is no host fallback.
11pub struct CudaHostPinnedBackend;
12
13impl CudaAllocOps for CudaHostPinnedBackend {
14    #[inline]
15    fn registry() -> &'static CudaAllocationRegistry {
16        &CUDA_HOST_PINNED_ALLOCATIONS
17    }
18
19    #[inline]
20    fn alloc_sym() -> *mut c_void {
21        loader::CU_MEM_HOST_ALLOC.load(Ordering::Acquire)
22    }
23
24    #[inline]
25    fn free_sym() -> *mut c_void {
26        loader::CU_MEM_FREE_HOST.load(Ordering::Acquire)
27    }
28
29    #[inline]
30    unsafe fn raw_alloc(alloc_sym: *mut c_void, size: usize) -> *mut u8 {
31        type CuMemHostAllocFn =
32            unsafe extern "system" fn(*mut *mut c_void, usize, u32) -> core::ffi::c_int;
33        // SAFETY: transmute maps the verified dynamic library symbol address
34        // to a function pointer with system calling convention.
35        let cu_mem_host_alloc: CuMemHostAllocFn = unsafe { core::mem::transmute(alloc_sym) };
36
37        let mut host_ptr: *mut c_void = core::ptr::null_mut();
38        // CU_MEMHOSTALLOC_DEVICEMAP = 0x02
39        // SAFETY: on a zero return, the driver wrote a host pointer valid for
40        // `size` bytes into `host_ptr`.
41        let res = unsafe { cu_mem_host_alloc(core::ptr::addr_of_mut!(host_ptr), size, 0x02) };
42        if res == 0 && !host_ptr.is_null() {
43            host_ptr as *mut u8
44        } else {
45            core::ptr::null_mut()
46        }
47    }
48
49    #[inline]
50    unsafe fn raw_free(free_sym: *mut c_void, ptr: *mut u8) -> core::ffi::c_int {
51        type CuMemFreeHostFn = unsafe extern "system" fn(*mut c_void) -> core::ffi::c_int;
52        // SAFETY: transmute maps the verified dynamic library symbol address
53        // to a function pointer with system calling convention.
54        let cu_mem_free_host: CuMemFreeHostFn = unsafe { core::mem::transmute(free_sym) };
55        // SAFETY: `ptr` is a live `cuMemHostAlloc` allocation per the caller
56        // contract.
57        unsafe { cu_mem_free_host(ptr as *mut c_void) }
58    }
59}
60
61impl_cuda_memory_backend!(CudaHostPinnedBackend);