use num_traits::Float;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
pub struct MemoryManager<T: Float> {
pools: HashMap<String, MemoryPool<T>>,
total_allocated: usize,
stats: MemoryStats,
}
pub struct MemoryPool<T: Float> {
available: Vec<Vec<T>>,
allocated_count: usize,
_buffer_size: usize,
_name: String,
}
#[derive(Debug, Clone)]
pub struct MemoryStats {
pub total_allocated: usize,
pub available: usize,
pub buffer_count: usize,
pub fragmentation_ratio: f64,
}
impl<T: Float> MemoryManager<T> {
pub fn new() -> Self {
Self {
pools: HashMap::new(),
total_allocated: 0,
stats: MemoryStats {
total_allocated: 0,
available: 0,
buffer_count: 0,
fragmentation_ratio: 0.0,
},
}
}
pub fn create_pool(&mut self, name: &str, buffer_size: usize) {
let pool = MemoryPool::new(name.to_string(), buffer_size);
self.pools.insert(name.to_string(), pool);
}
pub fn allocate(&mut self, pool_name: &str, size: usize) -> Result<Vec<T>, String> {
if let Some(pool) = self.pools.get_mut(pool_name) {
let buffer = pool.allocate(size)?;
self.total_allocated += size * std::mem::size_of::<T>();
self.update_stats();
Ok(buffer)
} else {
Err(format!("Pool '{pool_name}' not found"))
}
}
pub fn deallocate(&mut self, pool_name: &str, buffer: Vec<T>) -> Result<(), String> {
if let Some(pool) = self.pools.get_mut(pool_name) {
let size = buffer.len() * std::mem::size_of::<T>();
pool.deallocate(buffer);
self.total_allocated = self.total_allocated.saturating_sub(size);
self.update_stats();
Ok(())
} else {
Err(format!("Pool '{pool_name}' not found"))
}
}
pub fn get_stats(&self) -> MemoryStats {
self.stats.clone()
}
pub fn clear_all(&mut self) {
for pool in self.pools.values_mut() {
pool.clear();
}
self.total_allocated = 0;
self.update_stats();
}
fn update_stats(&mut self) {
let mut buffer_count = 0;
let mut available_buffers = 0;
for pool in self.pools.values() {
buffer_count += pool.allocated_count;
available_buffers += pool.available.len();
}
self.stats = MemoryStats {
total_allocated: self.total_allocated,
available: available_buffers * std::mem::size_of::<T>(),
buffer_count,
fragmentation_ratio: if buffer_count > 0 {
available_buffers as f64 / buffer_count as f64
} else {
0.0
},
};
}
}
impl<T: Float> Default for MemoryManager<T> {
fn default() -> Self {
Self::new()
}
}
impl<T: Float> MemoryPool<T> {
pub fn new(name: String, buffer_size: usize) -> Self {
Self {
available: Vec::new(),
allocated_count: 0,
_buffer_size: buffer_size,
_name: name,
}
}
pub fn allocate(&mut self, size: usize) -> Result<Vec<T>, String> {
if let Some(mut buffer) = self.available.pop() {
buffer.clear();
buffer.resize(size, T::zero());
self.allocated_count += 1;
Ok(buffer)
} else {
let buffer = vec![T::zero(); size];
self.allocated_count += 1;
Ok(buffer)
}
}
pub fn deallocate(&mut self, buffer: Vec<T>) {
self.available.push(buffer);
self.allocated_count = self.allocated_count.saturating_sub(1);
}
pub fn clear(&mut self) {
self.available.clear();
self.allocated_count = 0;
}
pub fn allocated_count(&self) -> usize {
self.allocated_count
}
pub fn available_count(&self) -> usize {
self.available.len()
}
}
lazy_static::lazy_static! {
static ref GLOBAL_MEMORY_MANAGER: Arc<Mutex<MemoryManager<f32>>> = Arc::new(Mutex::new(MemoryManager::new()));
}
pub fn get_global_memory_manager() -> Arc<Mutex<MemoryManager<f32>>> {
GLOBAL_MEMORY_MANAGER.clone()
}
pub fn init_default_pools() {
let mut manager = GLOBAL_MEMORY_MANAGER.lock().unwrap();
manager.create_pool("weights", 1024);
manager.create_pool("activations", 512);
manager.create_pool("gradients", 512);
manager.create_pool("temporary", 256);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_memory_manager_creation() {
let manager: MemoryManager<f32> = MemoryManager::new();
assert_eq!(manager.total_allocated, 0);
assert_eq!(manager.pools.len(), 0);
}
#[test]
fn test_pool_creation() {
let mut manager: MemoryManager<f32> = MemoryManager::new();
manager.create_pool("test", 100);
assert_eq!(manager.pools.len(), 1);
assert!(manager.pools.contains_key("test"));
}
#[test]
fn test_allocation_deallocation() {
let mut manager: MemoryManager<f32> = MemoryManager::new();
manager.create_pool("test", 100);
let buffer = manager.allocate("test", 50).unwrap();
assert_eq!(buffer.len(), 50);
assert!(manager.total_allocated > 0);
manager.deallocate("test", buffer).unwrap();
}
#[test]
fn test_memory_stats() {
let mut manager: MemoryManager<f32> = MemoryManager::new();
manager.create_pool("test", 100);
let stats = manager.get_stats();
assert_eq!(stats.buffer_count, 0);
assert_eq!(stats.total_allocated, 0);
let _buffer = manager.allocate("test", 50).unwrap();
let stats = manager.get_stats();
assert_eq!(stats.buffer_count, 1);
assert!(stats.total_allocated > 0);
}
#[test]
fn test_pool_reuse() {
let mut pool: MemoryPool<f32> = MemoryPool::new("test".to_string(), 100);
let buffer1 = pool.allocate(50).unwrap();
pool.deallocate(buffer1);
let buffer2 = pool.allocate(50).unwrap();
assert_eq!(buffer2.len(), 50);
assert_eq!(pool.available_count(), 0);
assert_eq!(pool.allocated_count(), 1);
}
}