use std::collections::HashSet;
use std::fmt;
use std::ptr::NonNull;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::sync::OnceLock;
use cudarc::driver::sys as cuda_sys;
use cudarc::driver::CudaDevice;
use parking_lot::Mutex;
use crate::advise::Advice;
use crate::unified::{DeviceId, UnifiedBacking, UnifiedError, UvmAdvice};
static CUDARC_FREE_FAILURES: AtomicU64 = AtomicU64::new(0);
static LEAKED_CUDA_ALLOCATIONS: OnceLock<Mutex<HashSet<u64>>> = OnceLock::new();
const MAX_TRACKED_FREE_FAILURES: usize = 4096;
static LEAK_TRACKING_OVERFLOWS: AtomicU64 = AtomicU64::new(0);
fn leaked_set() -> &'static Mutex<HashSet<u64>> {
LEAKED_CUDA_ALLOCATIONS.get_or_init(|| Mutex::new(HashSet::new()))
}
fn record_leak(ptr: u64) {
let mut set = leaked_set().lock();
if set.len() >= MAX_TRACKED_FREE_FAILURES && !set.contains(&ptr) {
LEAK_TRACKING_OVERFLOWS.fetch_add(1, Ordering::Relaxed);
return;
}
set.insert(ptr);
}
pub fn cudarc_free_failures() -> u64 {
CUDARC_FREE_FAILURES.load(Ordering::Relaxed)
}
pub fn cudarc_leak_tracking_overflows() -> u64 {
LEAK_TRACKING_OVERFLOWS.load(Ordering::Relaxed)
}
pub fn leaked_cuda_allocations() -> Vec<u64> {
leaked_set().lock().iter().copied().collect()
}
static DEVICE_CACHE: OnceLock<Mutex<std::collections::HashMap<u32, Arc<CudaDevice>>>> =
OnceLock::new();
fn device_cache() -> &'static Mutex<std::collections::HashMap<u32, Arc<CudaDevice>>> {
DEVICE_CACHE.get_or_init(|| Mutex::new(std::collections::HashMap::new()))
}
pub(crate) fn device_for(ordinal: u32) -> Result<Arc<CudaDevice>, UnifiedError> {
{
let cache = device_cache().lock();
if let Some(cached) = cache.get(&ordinal) {
return Ok(Arc::clone(cached));
}
}
let fresh = CudaDevice::new(ordinal as usize)
.map_err(|e| UnifiedError::Cuda(format!("CudaDevice::new({ordinal}): {e:?}")))?;
let mut cache = device_cache().lock();
let (canonical, loser) = match cache.entry(ordinal) {
std::collections::hash_map::Entry::Occupied(occupied) => {
(Arc::clone(occupied.get()), Some(fresh))
}
std::collections::hash_map::Entry::Vacant(vacant) => {
(Arc::clone(vacant.insert(fresh)), None)
}
};
drop(cache);
drop(loser);
Ok(canonical)
}
pub(crate) fn ensure_context_bound(device: &Arc<CudaDevice>) -> Result<(), UnifiedError> {
device
.bind_to_thread()
.map_err(|e| UnifiedError::Cuda(format!("CudaDevice::bind_to_thread: {e:?}")))
}
fn supports_managed_prefetch(ordinal: u32) -> bool {
unsafe {
let mut dev: cuda_sys::CUdevice = 0;
if cuda_sys::lib().cuDeviceGet(&mut dev as *mut cuda_sys::CUdevice, ordinal as i32)
!= cuda_sys::cudaError_enum::CUDA_SUCCESS
{
return false;
}
let mut val: i32 = 0;
if cuda_sys::lib().cuDeviceGetAttribute(
&mut val as *mut i32,
cuda_sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_CONCURRENT_MANAGED_ACCESS,
dev,
) != cuda_sys::cudaError_enum::CUDA_SUCCESS
{
return false;
}
val != 0
}
}
pub struct CudarcUnifiedBuffer {
ptr: NonNull<u8>,
size: usize,
device_id: DeviceId,
#[allow(dead_code)]
device: Arc<CudaDevice>,
}
unsafe impl Send for CudarcUnifiedBuffer {}
unsafe impl Sync for CudarcUnifiedBuffer {}
impl CudarcUnifiedBuffer {
pub fn new(size: usize) -> Result<Self, UnifiedError> {
Self::new_on(size, DeviceId::default())
}
pub fn new_on(size: usize, device_id: DeviceId) -> Result<Self, UnifiedError> {
if size == 0 {
return Err(UnifiedError::ZeroSize);
}
let device = device_for(device_id.0)?;
ensure_context_bound(&device)?;
let mut raw: cuda_sys::CUdeviceptr = 0;
const CU_MEM_ATTACH_GLOBAL: u32 = 1;
const _: () = assert!(
CU_MEM_ATTACH_GLOBAL == cuda_sys::CUmemAttach_flags::CU_MEM_ATTACH_GLOBAL as u32,
"cudarc renumbered CUmemAttach_flags::CU_MEM_ATTACH_GLOBAL; \
update the inlined constant in cudarc_backend.rs",
);
let res = unsafe {
cuda_sys::lib().cuMemAllocManaged(
&mut raw as *mut cuda_sys::CUdeviceptr,
size,
CU_MEM_ATTACH_GLOBAL,
)
};
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
return Err(UnifiedError::Cuda(format!("cuMemAllocManaged -> {res:?}")));
}
let ptr = NonNull::new(raw as *mut u8).ok_or_else(|| {
UnifiedError::Allocation("cuMemAllocManaged returned null with CUDA_SUCCESS".into())
})?;
Ok(Self {
ptr,
size,
device_id,
device,
})
}
pub fn len(&self) -> usize {
self.size
}
pub fn is_empty(&self) -> bool {
self.size == 0
}
pub fn as_ptr(&self) -> *const u8 {
self.ptr.as_ptr() as *const u8
}
pub fn as_mut_ptr(&mut self) -> *mut u8 {
self.ptr.as_ptr()
}
pub fn as_slice(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.size) }
}
pub fn as_mut_slice(&mut self) -> &mut [u8] {
unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.size) }
}
pub fn device_id(&self) -> DeviceId {
self.device_id
}
pub fn prefetch_to_device(&self) -> Result<(), UnifiedError> {
ensure_context_bound(&self.device)?;
if !supports_managed_prefetch(self.device_id.0) {
tracing::debug!(
target: "tensor_wasm_mem::cudarc_backend",
device = self.device_id.0,
"prefetch_to_device: skipped (device lacks CONCURRENT_MANAGED_ACCESS, e.g. Windows/WDDM)"
);
return Ok(());
}
let res = unsafe {
cuda_sys::lib().cuMemPrefetchAsync(
self.ptr.as_ptr() as cuda_sys::CUdeviceptr,
self.size,
self.device_id.0 as i32,
std::ptr::null_mut(),
)
};
if res == cuda_sys::cudaError_enum::CUDA_SUCCESS {
Ok(())
} else {
Err(UnifiedError::Cuda(format!(
"cuMemPrefetchAsync(device) -> {res:?}"
)))
}
}
pub fn prefetch_to_host(&self) -> Result<(), UnifiedError> {
const CU_DEVICE_CPU: i32 = -1;
ensure_context_bound(&self.device)?;
if !supports_managed_prefetch(self.device_id.0) {
tracing::debug!(
target: "tensor_wasm_mem::cudarc_backend",
device = self.device_id.0,
"prefetch_to_host: skipped (device lacks CONCURRENT_MANAGED_ACCESS, e.g. Windows/WDDM)"
);
return Ok(());
}
let res = unsafe {
cuda_sys::lib().cuMemPrefetchAsync(
self.ptr.as_ptr() as cuda_sys::CUdeviceptr,
self.size,
CU_DEVICE_CPU,
std::ptr::null_mut(),
)
};
if res == cuda_sys::cudaError_enum::CUDA_SUCCESS {
Ok(())
} else {
Err(UnifiedError::Cuda(format!(
"cuMemPrefetchAsync(host) -> {res:?}"
)))
}
}
}
impl fmt::Debug for CudarcUnifiedBuffer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CudarcUnifiedBuffer")
.field("ptr", &self.ptr.as_ptr())
.field("size", &self.size)
.field("device_id", &self.device_id)
.finish()
}
}
impl UnifiedBacking for CudarcUnifiedBuffer {
fn len(&self) -> usize {
CudarcUnifiedBuffer::len(self)
}
fn as_slice(&self) -> &[u8] {
CudarcUnifiedBuffer::as_slice(self)
}
fn as_mut_slice(&mut self) -> &mut [u8] {
CudarcUnifiedBuffer::as_mut_slice(self)
}
fn apply_advice(&self, hint: UvmAdvice) -> Result<(), UnifiedError> {
let advice = match hint {
UvmAdvice::SetReadMostly => Advice::ReadMostly,
UvmAdvice::UnsetReadMostly => {
return Err(UnifiedError::NotSupported {
feature: "apply_advice(UnsetReadMostly)",
backing: "cudarc",
});
}
UvmAdvice::SetPreferredLocation(d) => Advice::PreferredLocation(DeviceId(d)),
UvmAdvice::UnsetPreferredLocation => Advice::UnsetPreferredLocation,
UvmAdvice::SetAccessedBy(d) => Advice::AccessedBy(DeviceId(d)),
UvmAdvice::UnsetAccessedBy(d) => Advice::UnsetAccessedBy(DeviceId(d)),
};
self::apply_advice(self, advice)
}
fn prefetch_to_device(&self, device_ord: u32) -> Result<(), UnifiedError> {
if device_ord == self.device_id().0 {
CudarcUnifiedBuffer::prefetch_to_device(self)
} else {
Err(UnifiedError::NotSupported {
feature: "prefetch_to_device(non-owning-ordinal)",
backing: "cudarc",
})
}
}
fn prefetch_to_host(&self) -> Result<(), UnifiedError> {
CudarcUnifiedBuffer::prefetch_to_host(self)
}
}
impl Drop for CudarcUnifiedBuffer {
fn drop(&mut self) {
let raw_ptr = self.ptr.as_ptr();
let raw_ptr_u64 = raw_ptr as u64;
let res = unsafe { cuda_sys::lib().cuMemFree_v2(raw_ptr as cuda_sys::CUdeviceptr) };
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
let failures = CUDARC_FREE_FAILURES.fetch_add(1, Ordering::Relaxed) + 1;
record_leak(raw_ptr_u64);
tracing::error!(
target: "tensor_wasm_mem::cudarc_backend",
?res,
ptr = ?raw_ptr,
size = self.size,
total_failures = failures,
"cuMemFree_v2 failed in CudarcUnifiedBuffer::drop -- VA leaked, \
recorded in leaked_cuda_allocations() for operator audit (mem H4)",
);
self.ptr = NonNull::dangling();
}
}
}
pub fn apply_advice(buffer: &CudarcUnifiedBuffer, advice: Advice) -> Result<(), UnifiedError> {
ensure_context_bound(&buffer.device)?;
let ptr = buffer.as_ptr() as cuda_sys::CUdeviceptr;
let size = buffer.len();
let (advice_kind, device) = match advice {
Advice::ReadMostly => (
cuda_sys::CUmem_advise_enum::CU_MEM_ADVISE_SET_READ_MOSTLY,
0i32,
),
Advice::PreferredLocation(d) => (
cuda_sys::CUmem_advise_enum::CU_MEM_ADVISE_SET_PREFERRED_LOCATION,
d.0 as i32,
),
Advice::AccessedBy(d) => (
cuda_sys::CUmem_advise_enum::CU_MEM_ADVISE_SET_ACCESSED_BY,
d.0 as i32,
),
Advice::UnsetPreferredLocation => (
cuda_sys::CUmem_advise_enum::CU_MEM_ADVISE_UNSET_PREFERRED_LOCATION,
0i32,
),
Advice::UnsetAccessedBy(d) => (
cuda_sys::CUmem_advise_enum::CU_MEM_ADVISE_UNSET_ACCESSED_BY,
d.0 as i32,
),
};
let res = unsafe { cuda_sys::lib().cuMemAdvise(ptr, size, advice_kind, device) };
if res == cuda_sys::cudaError_enum::CUDA_SUCCESS {
Ok(())
} else {
Err(UnifiedError::Cuda(format!("cuMemAdvise -> {res:?}")))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn buffer_type_has_nonzero_size() {
assert!(std::mem::size_of::<CudarcUnifiedBuffer>() > 0);
}
#[test]
fn apply_advice_is_exported() {
let _f: fn(&CudarcUnifiedBuffer, Advice) -> Result<(), UnifiedError> = apply_advice;
}
#[test]
#[ignore = "requires CUDA hardware"]
fn allocate_and_drop_small_buffer() {
let mut b = CudarcUnifiedBuffer::new(64).expect("alloc");
assert_eq!(b.len(), 64);
b.as_mut_slice().copy_from_slice(&[0xAB; 64]);
assert!(b.as_slice().iter().all(|&v| v == 0xAB));
}
#[test]
#[ignore = "requires CUDA hardware"]
fn device_cache_returns_same_arc_for_same_ordinal() {
let a = device_for(0).expect("first device_for(0)");
let b = device_for(0).expect("second device_for(0)");
assert_eq!(
Arc::as_ptr(&a),
Arc::as_ptr(&b),
"device_for(0) must return clones of the cached Arc — \
a fresh `CudaDevice::new` on every call would re-arm the \
failed-construction race the T26 cache exists to close"
);
}
}