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,
}
#[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
);
Ok(Self {
device_id,
allocated_memory: AtomicU64::new(0),
max_memory,
config,
})
}
pub fn allocate(&mut self, size: usize) -> Result<CudaMemoryPtr, 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
));
}
self.allocated_memory.fetch_add(size_u64, Ordering::Relaxed);
debug!(
"Allocated {}MB CUDA memory on device {}, total allocated: {}MB",
size_u64 / 1024 / 1024,
self.device_id,
self.allocated_memory.load(Ordering::Relaxed) / 1024 / 1024
);
warn!(
"CUDA memory pool is using placeholder implementation. Actual memory allocation requires CudaDevice handle."
);
Ok(CudaMemoryPtr {
device_id: self.device_id,
size,
ptr: 0, })
}
pub fn deallocate(&mut self, ptr: CudaMemoryPtr) {
if ptr.size == 0 {
warn!("Attempted to deallocate CUDA memory with size 0");
return;
}
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
);
warn!(
"CUDA memory pool is using placeholder implementation. Actual memory deallocation requires CudaDevice handle."
);
}
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, Clone)]
pub struct CudaMemoryPtr {
pub device_id: i32,
pub size: usize,
pub ptr: usize,
}
#[cfg(feature = "cuda")]
impl Drop for CudaMemoryPtr {
fn drop(&mut self) {
debug!("Dropping CUDA memory ptr on device {}", self.device_id);
}
}
#[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 = CudaMemoryPool::new(0, config).unwrap();
let ptr = pool.allocate(512 * 1024 * 1024);
assert!(ptr.is_ok());
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 512 * 1024 * 1024);
let ptr = ptr.unwrap();
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 = CudaMemoryPool::new(0, config).unwrap();
let _ptr1 = pool.allocate(512 * 1024 * 1024).unwrap();
let ptr2 = pool.allocate(512 * 1024 * 1024);
assert!(ptr2.is_err());
}
#[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
);
}
}
}