use anyhow::{Result, anyhow};
use cudarc::driver::sys::{
self, CUmemAllocationType, CUmemLocationType, CUmemPool_attribute, CUmemPoolProps,
CUmemoryPool, CUresult, CUstream,
};
use cudarc::driver::{CudaContext, CudaStream};
use std::ptr;
use std::sync::{Arc, Mutex};
pub struct CudaMemPoolBuilder {
context: Arc<CudaContext>,
reserve_size: usize,
release_threshold: Option<u64>,
}
impl CudaMemPoolBuilder {
pub fn new(context: Arc<CudaContext>, reserve_size: usize) -> Self {
Self {
context,
reserve_size,
release_threshold: None,
}
}
pub fn release_threshold(mut self, threshold: u64) -> Self {
self.release_threshold = Some(threshold);
self
}
pub fn build(self) -> Result<CudaMemPool> {
let mut props: CUmemPoolProps = unsafe { std::mem::zeroed() };
props.allocType = CUmemAllocationType::CU_MEM_ALLOCATION_TYPE_PINNED;
props.location.type_ = CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE;
props.location.id = self.context.cu_device();
let mut pool: CUmemoryPool = ptr::null_mut();
let result = unsafe { sys::cuMemPoolCreate(&mut pool, &props) };
if result != CUresult::CUDA_SUCCESS {
return Err(anyhow!("cuMemPoolCreate failed with error: {:?}", result));
}
if let Some(threshold) = self.release_threshold {
let result = unsafe {
sys::cuMemPoolSetAttribute(
pool,
CUmemPool_attribute::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
&threshold as *const u64 as *mut std::ffi::c_void,
)
};
if result != CUresult::CUDA_SUCCESS {
unsafe { sys::cuMemPoolDestroy(pool) };
return Err(anyhow!(
"cuMemPoolSetAttribute failed with error: {:?}",
result
));
}
}
let cuda_pool = CudaMemPool {
inner: Mutex::new(pool),
};
if self.reserve_size > 0 {
let stream = self.context.new_stream()?;
let ptr = cuda_pool.alloc_async(self.reserve_size, &stream)?;
cuda_pool.free_async(ptr, &stream)?;
let result = unsafe { sys::cuStreamSynchronize(stream.cu_stream()) };
if result != CUresult::CUDA_SUCCESS {
return Err(anyhow!(
"cuStreamSynchronize failed with error: {:?}",
result
));
}
}
Ok(cuda_pool)
}
}
pub struct CudaMemPool {
inner: Mutex<CUmemoryPool>,
}
unsafe impl Send for CudaMemPool {}
unsafe impl Sync for CudaMemPool {}
impl CudaMemPool {
pub fn builder(context: Arc<CudaContext>, reserve_size: usize) -> CudaMemPoolBuilder {
CudaMemPoolBuilder::new(context, reserve_size)
}
pub fn alloc_async(&self, size: usize, stream: &CudaStream) -> Result<u64> {
unsafe { self.alloc_async_raw(size, stream.cu_stream()) }
}
pub unsafe fn alloc_async_raw(&self, size: usize, stream: CUstream) -> Result<u64> {
let pool = self
.inner
.lock()
.map_err(|e| anyhow!("mutex poisoned: {}", e))?;
let mut ptr: u64 = 0;
let result = unsafe { sys::cuMemAllocFromPoolAsync(&mut ptr, size, *pool, stream) };
if result != CUresult::CUDA_SUCCESS {
return Err(anyhow!(
"cuMemAllocFromPoolAsync failed with error: {:?}",
result
));
}
Ok(ptr)
}
pub fn free_async(&self, ptr: u64, stream: &CudaStream) -> Result<()> {
unsafe { self.free_async_raw(ptr, stream.cu_stream()) }
}
pub unsafe fn free_async_raw(&self, ptr: u64, stream: CUstream) -> Result<()> {
let result = unsafe { sys::cuMemFreeAsync(ptr, stream) };
if result != CUresult::CUDA_SUCCESS {
return Err(anyhow!("cuMemFreeAsync failed with error: {:?}", result));
}
Ok(())
}
}
impl Drop for CudaMemPool {
fn drop(&mut self) {
let pool = self
.inner
.get_mut()
.expect("mutex should not be poisoned during drop");
let result = unsafe { sys::cuMemPoolDestroy(*pool) };
if result != CUresult::CUDA_SUCCESS {
tracing::warn!("cuMemPoolDestroy failed with error: {:?}", result);
}
}
}
#[cfg(all(test, feature = "testing-cuda"))]
mod tests {
use super::*;
#[test]
fn test_pool_creation_with_builder() {
let context = match CudaContext::new(0) {
Ok(ctx) => ctx,
Err(e) => {
eprintln!("Skipping test - no CUDA device: {:?}", e);
return;
}
};
let result = CudaMemPool::builder(context.clone(), 1024 * 1024) .release_threshold(64 * 1024 * 1024) .build();
if result.is_err() {
eprintln!("Skipping test - pool creation failed: {:?}", result.err());
return;
}
let pool = result.unwrap();
drop(pool);
}
#[test]
fn test_pool_creation_no_threshold() {
let context = match CudaContext::new(0) {
Ok(ctx) => ctx,
Err(e) => {
eprintln!("Skipping test - no CUDA device: {:?}", e);
return;
}
};
let result = CudaMemPool::builder(context, 0).build();
if result.is_err() {
eprintln!("Skipping test - pool creation failed: {:?}", result.err());
return;
}
let pool = result.unwrap();
drop(pool);
}
}