Skip to main content

mnemosyne_backend/backends/cuda/
device.rs

1use core::ffi::c_void;
2use core::sync::atomic::Ordering;
3
4use mnemosyne_core::MemoryBackend;
5
6use super::registry::CUDA_DEVICE_ALLOCATIONS;
7use super::{CudaAllocOps, CudaAllocationRegistry, loader, managed_raw_alloc, managed_raw_free};
8
9/// A memory backend allocating CUDA device memory.
10///
11/// Under the hood, this uses CUDA unified memory (`cuMemAllocManaged`) and
12/// advises the driver to prefer device placement (`cuMemAdvise` with
13/// `CU_MEM_ADVISE_SET_PREFERRED_LOCATION`). This allows the host CPU to write
14/// allocator metadata in-band without segfaulting, while keeping the
15/// allocation device-preferred for optimal kernel performance.
16///
17/// `allocate` returns null on failure (driver unavailable, driver allocation
18/// failure, or registry full); there is no host fallback.
19pub struct CudaDeviceBackend;
20
21impl CudaAllocOps for CudaDeviceBackend {
22    #[inline]
23    fn registry() -> &'static CudaAllocationRegistry {
24        &CUDA_DEVICE_ALLOCATIONS
25    }
26
27    #[inline]
28    fn alloc_sym() -> *mut c_void {
29        loader::CU_MEM_ALLOC_MANAGED.load(Ordering::Acquire)
30    }
31
32    #[inline]
33    fn free_sym() -> *mut c_void {
34        loader::CU_MEM_FREE.load(Ordering::Acquire)
35    }
36
37    #[inline]
38    unsafe fn raw_alloc(alloc_sym: *mut c_void, size: usize) -> *mut u8 {
39        // SAFETY: forwarded caller contract (resolved `cuMemAllocManaged`).
40        unsafe { managed_raw_alloc(alloc_sym, size) }
41    }
42
43    #[inline]
44    unsafe fn raw_free(free_sym: *mut c_void, ptr: *mut u8) -> core::ffi::c_int {
45        // SAFETY: forwarded caller contract (resolved `cuMemFree`, live ptr).
46        unsafe { managed_raw_free(free_sym, ptr) }
47    }
48
49    #[inline]
50    unsafe fn post_register(ptr: *mut u8, size: usize) {
51        let advise_sym = loader::CU_MEM_ADVISE.load(Ordering::Acquire);
52        if advise_sym.is_null() {
53            return;
54        }
55        type CuMemAdviseFn = unsafe extern "system" fn(u64, usize, u32, i32) -> core::ffi::c_int;
56        // SAFETY: transmute maps the verified dynamic library symbol address
57        // to a function pointer with system calling convention.
58        let cu_mem_advise: CuMemAdviseFn = unsafe { core::mem::transmute(advise_sym) };
59        // CU_MEM_ADVISE_SET_PREFERRED_LOCATION = 3, device ordinal 0.
60        // Placement advice is best-effort tuning: a nonzero status leaves the
61        // allocation valid and host-accessible, so there is no failure to
62        // surface or recover from here.
63        // SAFETY: `ptr` is a live managed allocation of `size` bytes per the
64        // trait contract.
65        let _advise_status = unsafe { cu_mem_advise(ptr as u64, size, 3, 0) };
66    }
67}
68
69impl_cuda_memory_backend!(CudaDeviceBackend);
70
71/// Allocates through the shared device driver while allowing the arena to
72/// keep a distinct pool identity for a device memory tier.
73///
74/// CUDA's managed allocation API does not expose an HBM-versus-GDDR selector;
75/// the tier-specific wrappers therefore preserve the provider's device
76/// allocation semantics and split allocator-local retention state. A future
77/// provider with an explicit tier selector can replace the wrapper's
78/// monomorphic forwarding without changing the heap dispatch surface.
79#[inline(always)]
80unsafe fn allocate_device_tier(size: usize) -> *mut u8 {
81    // SAFETY: the caller upholds the `MemoryBackend::allocate` contract, and
82    // `CudaDeviceBackend` is the shared driver-backed implementation.
83    unsafe { CudaDeviceBackend::allocate(size) }
84}
85
86/// Releases a tier-keyed device allocation through the shared CUDA driver.
87#[inline(always)]
88unsafe fn deallocate_device_tier(ptr: *mut u8, size: usize) -> bool {
89    // SAFETY: the caller upholds the `MemoryBackend::deallocate` contract, and
90    // the pointer was allocated by the shared device backend.
91    unsafe { CudaDeviceBackend::deallocate(ptr, size) }
92}
93
94/// A zero-sized device backend with an allocator-local HBM pool identity.
95///
96/// The current CUDA provider uses the same managed-memory driver operation as
97/// [`CudaDeviceBackend`]. This type separates segment and thread-local pool
98/// ownership for `MemoryTier::Hbm` without adding a runtime dispatch branch.
99pub struct CudaHbmBackend;
100
101/// A zero-sized device backend with an allocator-local GDDR pool identity.
102///
103/// The current CUDA provider uses the same managed-memory driver operation as
104/// [`CudaDeviceBackend`]. This type separates segment and thread-local pool
105/// ownership for `MemoryTier::Gddr` without adding a runtime dispatch branch.
106pub struct CudaGddrBackend;
107
108const _: () = assert!(
109    core::mem::size_of::<CudaHbmBackend>() == 0
110        && core::mem::size_of::<CudaGddrBackend>() == 0
111        && core::mem::align_of::<CudaHbmBackend>() == 1
112        && core::mem::align_of::<CudaGddrBackend>() == 1
113);
114
115impl_device_tier_backend!(CudaHbmBackend);
116impl_device_tier_backend!(CudaGddrBackend);