use log::{debug, info, warn};
use std::sync::atomic::{AtomicU64, Ordering};
use super::config::CudaPoolConfig;
#[cfg(feature = "cuda")]
pub struct CudaMemoryPool {
device_id: i32,
allocated_memory: AtomicU64,
max_memory: u64,
config: CudaPoolConfig,
_ctx: cudarc::driver::sys::CUcontext,
}
#[cfg(feature = "cuda")]
impl CudaMemoryPool {
pub fn new(device_id: i32, config: CudaPoolConfig) -> Result<Self, String> {
let max_memory = (config.max_memory_mb * 1024 * 1024) as u64;
info!(
"Creating CudaMemoryPool for device {} with max_memory={}MB",
device_id, config.max_memory_mb
);
cudarc::driver::result::init().map_err(|e| format!("CUDA driver init failed: {}", e))?;
let device = cudarc::driver::result::device::get(device_id)
.map_err(|e| format!("CUDA device {} not found: {}", device_id, e))?;
let ctx = unsafe {
cudarc::driver::result::primary_ctx::retain(device).map_err(|e| {
format!(
"CUDA context creation failed for device {}: {}",
device_id, e
)
})?
};
info!("CUDA context created for device {}", device_id);
Ok(Self {
device_id,
allocated_memory: AtomicU64::new(0),
max_memory,
config,
_ctx: ctx,
})
}
pub fn allocate(&mut self, size: usize) -> Result<CudaMemoryPtr, String> {
if size == 0 {
return Err("CUDA allocation size must be > 0".to_string());
}
let size_u64 = size as u64;
let current_allocated = self.allocated_memory.load(Ordering::Relaxed);
let available = self.max_memory.saturating_sub(current_allocated);
if size_u64 > available {
return Err(format!(
"Insufficient CUDA memory: need {}MB, available {}MB",
size_u64 / 1024 / 1024,
available / 1024 / 1024
));
}
let dev_ptr = unsafe {
cudarc::driver::result::malloc_sync(size)
.map_err(|e| format!("CUDA malloc failed for {} bytes: {}", size, e))?
};
self.allocated_memory.fetch_add(size_u64, Ordering::Relaxed);
debug!(
"Allocated {}MB CUDA memory on device {} at {:?}, total allocated: {}MB",
size_u64 / 1024 / 1024,
self.device_id,
dev_ptr,
self.allocated_memory.load(Ordering::Relaxed) / 1024 / 1024
);
Ok(CudaMemoryPtr {
device_id: self.device_id,
size,
ptr: dev_ptr,
})
}
pub fn deallocate(&mut self, ptr: CudaMemoryPtr) {
if ptr.size == 0 {
warn!("Attempted to deallocate CUDA memory with size 0");
return;
}
if let Err(e) = unsafe { cudarc::driver::result::free_sync(ptr.ptr) } {
warn!("CUDA free failed for device {}: {}", self.device_id, e);
}
self.allocated_memory
.fetch_sub(ptr.size as u64, Ordering::Relaxed);
debug!(
"Deallocated {}MB CUDA memory on device {}, total allocated: {}MB",
ptr.size / 1024 / 1024,
self.device_id,
self.allocated_memory.load(Ordering::Relaxed) / 1024 / 1024
);
std::mem::forget(ptr);
}
pub fn get_memory_usage(&self) -> (u64, u64) {
let used = self.allocated_memory.load(Ordering::Relaxed);
(used, self.max_memory)
}
pub fn get_memory_usage_percent(&self) -> f64 {
let used = self.allocated_memory.load(Ordering::Relaxed) as f64;
let total = self.max_memory as f64;
(used / total) * 100.0
}
pub fn clear(&mut self) {
info!("Clearing CUDA memory pool on device {}...", self.device_id);
self.allocated_memory.store(0, Ordering::Relaxed);
info!("CUDA memory pool cleared");
}
}
#[cfg(feature = "cuda")]
#[derive(Debug)]
pub struct CudaMemoryPtr {
pub device_id: i32,
pub size: usize,
pub ptr: cudarc::driver::sys::CUdeviceptr,
}
#[cfg(feature = "cuda")]
impl Drop for CudaMemoryPtr {
fn drop(&mut self) {
if self.size > 0 && self.ptr != cudarc::driver::sys::CUdeviceptr::default() {
debug!(
"Dropping CUDA memory ptr on device {} ({} bytes) — freeing via cuMemFree_v2",
self.device_id, self.size
);
if let Err(e) = unsafe { cudarc::driver::result::free_sync(self.ptr) } {
warn!(
"CUDA free on drop failed for device {}: {}",
self.device_id, e
);
}
}
}
}
#[cfg(not(feature = "cuda"))]
pub struct CudaMemoryPool {
_device_id: i32,
_config: CudaPoolConfig,
}
#[cfg(not(feature = "cuda"))]
impl CudaMemoryPool {
pub fn new(_device_id: i32, _config: CudaPoolConfig) -> Result<Self, String> {
warn!("CUDA feature not enabled, CudaMemoryPool will be no-op");
Ok(Self {
_device_id: 0,
_config,
})
}
pub fn allocate(&mut self, _size: usize) -> Result<CudaMemoryPtr, String> {
Err("CUDA feature not enabled".to_string())
}
pub fn deallocate(&mut self, _ptr: CudaMemoryPtr) {
}
pub fn get_memory_usage(&self) -> (u64, u64) {
(0, 0)
}
pub fn get_memory_usage_percent(&self) -> f64 {
0.0
}
pub fn clear(&mut self) {
}
}
#[cfg(not(feature = "cuda"))]
#[derive(Debug, Clone)]
pub struct CudaMemoryPtr {
_size: usize,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_pool_creation() {
let config = CudaPoolConfig::default();
let pool = CudaMemoryPool::new(0, config);
assert!(pool.is_ok());
let pool = pool.unwrap();
let (used, total) = pool.get_memory_usage();
assert_eq!(used, 0);
assert!(total > 0);
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_allocate_deallocate() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return, };
let ptr = match pool.allocate(512 * 1024 * 1024) {
Ok(p) => p,
Err(_) => return, };
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 512 * 1024 * 1024);
pool.deallocate(ptr);
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 0);
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_insufficient_memory() {
let config = CudaPoolConfig {
max_memory_mb: 512,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
let _ptr1 = match pool.allocate(512 * 1024 * 1024) {
Ok(p) => p,
Err(_) => return,
};
let ptr2 = pool.allocate(512 * 1024 * 1024);
assert!(ptr2.is_err());
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_allocate_zero_size_returns_error() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
let result = pool.allocate(0);
assert!(result.is_err());
assert!(result.unwrap_err().contains("size must be > 0"));
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_memory_usage_percent() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
assert_eq!(pool.get_memory_usage_percent(), 0.0);
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_pool_clear() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
let _ptr = match pool.allocate(1024 * 1024) {
Ok(p) => p,
Err(_) => return,
};
let (used_before, _) = pool.get_memory_usage();
assert!(used_before > 0);
pool.clear();
let (used_after, _) = pool.get_memory_usage();
assert_eq!(used_after, 0);
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_multiple_allocations() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
let ptr1 = match pool.allocate(64 * 1024 * 1024) {
Ok(p) => p,
Err(_) => return,
};
let ptr2 = pool.allocate(64 * 1024 * 1024);
assert!(ptr2.is_ok());
let (used, total) = pool.get_memory_usage();
assert_eq!(used, 128 * 1024 * 1024);
assert_eq!(total, 1024 * 1024 * 1024);
pool.deallocate(ptr1);
let ptr3 = pool.allocate(32 * 1024 * 1024);
assert!(ptr3.is_ok());
}
#[test]
#[cfg(feature = "cuda")]
fn test_cuda_memory_ptr_drop_frees() {
let config = CudaPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = match CudaMemoryPool::new(0, config) {
Ok(p) => p,
Err(_) => return,
};
{
let _ptr = match pool.allocate(1024 * 1024) {
Ok(p) => p,
Err(_) => return,
};
}
}
#[test]
#[cfg(not(feature = "cuda"))]
fn test_cuda_pool_no_cuda() {
let config = CudaPoolConfig::default();
let pool = CudaMemoryPool::new(0, config);
assert!(pool.is_ok());
let mut pool = pool.unwrap();
let ptr = pool.allocate(1024);
assert!(ptr.is_err());
}
#[test]
#[cfg(not(feature = "cuda"))]
fn test_cuda_pool_no_cuda_memory_usage() {
let config = CudaPoolConfig {
max_memory_mb: 2048,
..Default::default()
};
let pool = CudaMemoryPool::new(0, config).unwrap();
let (used, total) = pool.get_memory_usage();
assert_eq!(used, 0);
assert_eq!(total, 0);
let percent = pool.get_memory_usage_percent();
assert_eq!(percent, 0.0);
}
#[test]
#[cfg(not(feature = "cuda"))]
fn test_cuda_pool_no_cuda_clear_noop() {
let config = CudaPoolConfig::default();
let mut pool = CudaMemoryPool::new(1, config).unwrap();
pool.clear();
pool.clear();
let (used, total) = pool.get_memory_usage();
assert_eq!(used, 0);
assert_eq!(total, 0);
}
#[test]
#[cfg(not(feature = "cuda"))]
fn test_cuda_pool_no_cuda_new_with_different_device_ids() {
for device_id in [0, 1, -1, 42] {
let config = CudaPoolConfig::default();
let pool = CudaMemoryPool::new(device_id, config);
assert!(
pool.is_ok(),
"Failed to create pool for device {}",
device_id
);
}
}
}