#![cfg(feature = "cudarc-backend")]
use std::ptr::NonNull;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use cudarc::driver::sys as cuda_sys;
use cudarc::driver::CudaDevice;
use tensor_wasm_core::mem_pool::DriverMemPool;
use crate::cudarc_backend::{device_for, ensure_context_bound};
use crate::unified::UnifiedError;
#[derive(Debug)]
pub struct TenantMemPool {
pool: cuda_sys::CUmemoryPool,
cap_bytes: AtomicU64,
live_bytes: AtomicU64,
device_ordinal: u32,
#[allow(dead_code)]
device: Arc<CudaDevice>,
}
impl TenantMemPool {
pub fn new(device_ordinal: u32, cap_bytes: u64) -> Result<Self, MemPoolError> {
let device = device_for(device_ordinal)
.map_err(|e| MemPoolError::Device(format!("device_for({device_ordinal}): {e:?}")))?;
ensure_context_bound(&device)
.map_err(|e| MemPoolError::Device(format!("ensure_context_bound: {e:?}")))?;
unsafe {
let mut pool: cuda_sys::CUmemoryPool = std::ptr::null_mut();
let mut props: cuda_sys::CUmemPoolProps = std::mem::zeroed();
props.allocType = cuda_sys::CUmemAllocationType_enum::CU_MEM_ALLOCATION_TYPE_PINNED;
props.handleTypes = cuda_sys::CUmemAllocationHandleType_enum::CU_MEM_HANDLE_TYPE_NONE;
props.location.type_ = cuda_sys::CUmemLocationType_enum::CU_MEM_LOCATION_TYPE_DEVICE;
props.location.id = device_ordinal as core::ffi::c_int;
let res =
cuda_sys::lib().cuMemPoolCreate(&mut pool as *mut cuda_sys::CUmemoryPool, &props);
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
return Err(MemPoolError::Create(format!("{res:?}")));
}
let res = cuda_sys::lib().cuMemPoolSetAttribute(
pool,
cuda_sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
&cap_bytes as *const u64 as *mut core::ffi::c_void,
);
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
let _ = cuda_sys::lib().cuMemPoolDestroy(pool);
return Err(MemPoolError::SetAttribute(format!("{res:?}")));
}
Ok(Self {
pool,
cap_bytes: AtomicU64::new(cap_bytes),
live_bytes: AtomicU64::new(0),
device_ordinal,
device,
})
}
}
pub fn new_on_default_device(cap_bytes: u64) -> Result<Self, MemPoolError> {
Self::new(0, cap_bytes)
}
pub fn cap_bytes(&self) -> u64 {
self.cap_bytes.load(Ordering::Relaxed)
}
pub fn device_ordinal(&self) -> u32 {
self.device_ordinal
}
pub fn raw_handle(&self) -> cuda_sys::CUmemoryPool {
self.pool
}
pub(crate) fn allocate(&self, size: usize) -> Result<NonNull<u8>, UnifiedError> {
let cap = self.cap_bytes.load(Ordering::Acquire);
let size_u64 = size as u64;
let mut current = self.live_bytes.load(Ordering::Acquire);
loop {
let next = reserve_step(cap, current, size_u64)?;
match self.live_bytes.compare_exchange_weak(
current,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
if let Err(e) = ensure_context_bound(&self.device) {
self.release_bytes(size); return Err(e);
}
let mut raw: cuda_sys::CUdeviceptr = 0;
let res = unsafe {
cuda_sys::lib().cuMemAllocFromPoolAsync(
&mut raw as *mut cuda_sys::CUdeviceptr,
size,
self.pool,
std::ptr::null_mut(),
)
};
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
self.release_bytes(size); return Err(UnifiedError::Cuda(format!(
"cuMemAllocFromPoolAsync -> {res:?}"
)));
}
NonNull::new(raw as *mut u8).ok_or_else(|| {
self.release_bytes(size); UnifiedError::Allocation(
"cuMemAllocFromPoolAsync returned null with CUDA_SUCCESS".into(),
)
})
}
pub(crate) fn release_bytes(&self, size: usize) {
let size_u64 = size as u64;
let _ = self
.live_bytes
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
Some(cur.saturating_sub(size_u64))
});
}
pub fn live_bytes(&self) -> u64 {
self.live_bytes.load(Ordering::Acquire)
}
pub(crate) fn deallocate(&self, ptr: NonNull<u8>) -> Result<(), UnifiedError> {
ensure_context_bound(&self.device)?;
let res = unsafe {
cuda_sys::lib()
.cuMemFreeAsync(ptr.as_ptr() as cuda_sys::CUdeviceptr, std::ptr::null_mut())
};
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
return Err(UnifiedError::Cuda(format!("cuMemFreeAsync -> {res:?}")));
}
Ok(())
}
}
fn reserve_step(cap: u64, current: u64, size: u64) -> Result<u64, UnifiedError> {
let next = current.checked_add(size).ok_or_else(|| {
UnifiedError::Cuda(format!(
"CUDA_ERROR_OUT_OF_MEMORY: per-tenant pool reservation overflow \
(live={current}, requested={size})"
))
})?;
if next > cap {
return Err(UnifiedError::Cuda(format!(
"CUDA_ERROR_OUT_OF_MEMORY: allocation of {size} bytes would exceed the \
per-tenant GPU memory cap (cap={cap}, live={current}). The cap is \
enforced host-side; CU_MEMPOOL_ATTR_RELEASE_THRESHOLD is only a \
retention hint, not an allocation ceiling."
)));
}
Ok(next)
}
impl DriverMemPool for TenantMemPool {
fn set_release_threshold(&self, bytes: u64) -> Result<(), MemPoolError> {
ensure_context_bound(&self.device)
.map_err(|e| MemPoolError::Device(format!("ensure_context_bound: {e:?}")))?;
let res = unsafe {
cuda_sys::lib().cuMemPoolSetAttribute(
self.pool,
cuda_sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
&bytes as *const u64 as *mut core::ffi::c_void,
)
};
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
return Err(MemPoolError::SetAttribute(format!("{res:?}")));
}
self.cap_bytes.store(bytes, Ordering::Relaxed);
Ok(())
}
fn release_threshold(&self) -> Option<u64> {
Some(self.cap_bytes.load(Ordering::Relaxed))
}
}
impl Drop for TenantMemPool {
fn drop(&mut self) {
if self.pool.is_null() {
return;
}
let res = unsafe { cuda_sys::lib().cuMemPoolDestroy(self.pool) };
if res != cuda_sys::cudaError_enum::CUDA_SUCCESS {
tracing::error!(
target: "tensor_wasm_mem::cuda_mem_pool",
?res,
cap_bytes = self.cap_bytes.load(Ordering::Relaxed),
device_ordinal = self.device_ordinal,
"cuMemPoolDestroy failed in TenantMemPool::drop",
);
}
self.pool = std::ptr::null_mut();
}
}
unsafe impl Send for TenantMemPool {}
unsafe impl Sync for TenantMemPool {}
pub use tensor_wasm_core::mem_pool::MemPoolError;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tenant_mem_pool_type_has_nonzero_size() {
assert!(std::mem::size_of::<TenantMemPool>() > 0);
}
#[test]
fn tenant_mem_pool_is_a_driver_mem_pool() {
fn assert_driver_mem_pool<T: DriverMemPool>() {}
assert_driver_mem_pool::<TenantMemPool>();
fn _accepts_dyn(_: Arc<dyn DriverMemPool>) {}
}
#[test]
fn mem_pool_error_display_non_empty() {
let e = MemPoolError::Create("CUDA_ERROR_OUT_OF_MEMORY".into());
assert!(format!("{e}").contains("cuMemPoolCreate failed"));
let e = MemPoolError::SetAttribute("CUDA_ERROR_INVALID_VALUE".into());
assert!(format!("{e}").contains("cuMemPoolSetAttribute failed"));
let e = MemPoolError::NotInitialized;
assert!(format!("{e}").contains("not initialized"));
let e = MemPoolError::Device("device_for(7): CudaDevice::new(7): ...".into());
assert!(format!("{e}").contains("device retain failed"));
}
#[test]
fn reserve_step_admits_up_to_the_cap() {
const CAP: u64 = 64 * 1024 * 1024;
assert_eq!(reserve_step(CAP, 0, CAP).unwrap(), CAP);
assert_eq!(
reserve_step(CAP, 16 * 1024 * 1024, 16 * 1024 * 1024).unwrap(),
32 * 1024 * 1024
);
assert_eq!(
reserve_step(CAP, 48 * 1024 * 1024, 16 * 1024 * 1024).unwrap(),
CAP
);
}
#[test]
fn reserve_step_rejects_over_cap_with_oom_shape() {
const CAP: u64 = 64 * 1024 * 1024;
let err = reserve_step(CAP, 0, 128 * 1024 * 1024).unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("OUT_OF_MEMORY"),
"over-cap rejection must be OOM-shaped so callers match the driver \
error; got: {msg}"
);
assert!(reserve_step(CAP, CAP, 1).is_err());
assert_eq!(reserve_step(CAP, CAP, 0).unwrap(), CAP);
}
#[test]
fn reserve_step_rejects_overflow_as_oom() {
let err = reserve_step(u64::MAX, u64::MAX - 4, 16).unwrap_err();
assert!(format!("{err}").contains("OUT_OF_MEMORY"));
}
}