use log::{debug, info, warn};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Instant;
use super::config::ModelWeightPoolConfig;
pub struct ModelWeightPool {
device_id: String,
allocated_memory: AtomicU64,
max_memory: u64,
model_slots: HashMap<String, ModelSlot>,
config: ModelWeightPoolConfig,
cache_enabled: bool,
}
#[derive(Debug, Clone)]
pub struct ModelSlot {
pub model_name: String,
pub memory_allocated: u64,
pub is_loaded: bool,
pub last_used: Instant,
pub loaded_at: Option<Instant>,
}
impl ModelWeightPool {
pub fn new(device_id: String, config: ModelWeightPoolConfig) -> Self {
let max_memory = (config.max_memory_mb * 1024 * 1024) as u64;
info!(
"Creating ModelWeightPool for device {} with max_memory={}MB",
device_id, config.max_memory_mb
);
let cache_enabled = config.cache_models;
Self {
device_id,
allocated_memory: AtomicU64::new(0),
max_memory,
model_slots: HashMap::new(),
config,
cache_enabled,
}
}
pub fn can_load_model(&self, memory_bytes: u64) -> bool {
let current_allocated = self.allocated_memory.load(Ordering::Relaxed);
let available = self.max_memory.saturating_sub(current_allocated);
if memory_bytes <= available {
return true;
}
if self.cache_enabled {
let reclaimable = self.calculate_reclaimable_memory();
if memory_bytes <= available + reclaimable {
return true;
}
}
false
}
pub fn allocate_for_model(
&mut self,
model_name: &str,
memory_bytes: u64,
) -> Result<(), String> {
if self.model_slots.contains_key(model_name) {
warn!("Model {} already allocated", model_name);
return Ok(());
}
loop {
let current_allocated = self.allocated_memory.load(Ordering::Acquire);
let available = self.max_memory.saturating_sub(current_allocated);
if memory_bytes > available {
if self.cache_enabled {
let needed = memory_bytes - available;
let freed = self.reclaim_memory(needed);
if freed < needed {
return Err(format!(
"Insufficient memory: need {}MB, available {}MB, freed {}MB",
memory_bytes / 1024 / 1024,
available / 1024 / 1024,
freed / 1024 / 1024
));
}
} else {
return Err(format!(
"Insufficient memory: need {}MB, available {}MB",
memory_bytes / 1024 / 1024,
available / 1024 / 1024
));
}
}
let new_allocated = current_allocated + memory_bytes;
match self.allocated_memory.compare_exchange_weak(
current_allocated,
new_allocated,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
break;
}
Err(_) => {
debug!("Memory allocation race detected, retrying...");
continue;
}
}
}
let slot = ModelSlot {
model_name: model_name.to_string(),
memory_allocated: memory_bytes,
is_loaded: true,
last_used: Instant::now(),
loaded_at: Some(Instant::now()),
};
self.model_slots.insert(model_name.to_string(), slot);
info!(
"Allocated {}MB for model {}, total allocated: {}MB",
memory_bytes / 1024 / 1024,
model_name,
self.allocated_memory.load(Ordering::Relaxed) / 1024 / 1024
);
Ok(())
}
pub fn release_model(&mut self, model_name: &str) {
if let Some(slot) = self.model_slots.remove(model_name) {
self.allocated_memory
.fetch_sub(slot.memory_allocated, Ordering::Relaxed);
info!(
"Released {}MB for model {}, total allocated: {}MB",
slot.memory_allocated / 1024 / 1024,
model_name,
self.allocated_memory.load(Ordering::Relaxed) / 1024 / 1024
);
}
}
pub fn update_model_usage(&mut self, model_name: &str) {
if let Some(slot) = self.model_slots.get_mut(model_name) {
slot.last_used = Instant::now();
slot.is_loaded = true;
debug!("Updated usage time for model {}", model_name);
}
}
pub fn mark_model_unloaded(&mut self, model_name: &str) {
if let Some(slot) = self.model_slots.get_mut(model_name) {
slot.is_loaded = false;
debug!("Marked model {} as unloaded", model_name);
}
}
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 get_model_slots(&self) -> Vec<ModelSlot> {
self.model_slots.values().cloned().collect()
}
pub fn get_model_slot(&self, model_name: &str) -> Option<&ModelSlot> {
self.model_slots.get(model_name)
}
fn calculate_reclaimable_memory(&self) -> u64 {
let mut reclaimable = 0u64;
for slot in self.model_slots.values() {
if !slot.is_loaded {
reclaimable += slot.memory_allocated;
}
}
reclaimable
}
fn reclaim_memory(&mut self, needed: u64) -> u64 {
let mut freed = 0u64;
let mut slots: Vec<_> = self
.model_slots
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
slots.sort_by_key(|a| a.1.last_used);
for (model_name, slot) in slots {
if !slot.is_loaded {
self.allocated_memory
.fetch_sub(slot.memory_allocated, Ordering::Relaxed);
self.model_slots.remove(&model_name);
freed += slot.memory_allocated;
debug!(
"Reclaimed {}MB from model {}",
slot.memory_allocated / 1024 / 1024,
model_name
);
if freed >= needed {
break;
}
}
}
freed
}
pub fn clear(&mut self) {
info!("Clearing model weight pool...");
self.model_slots.clear();
self.allocated_memory.store(0, Ordering::Relaxed);
info!("Model weight pool cleared");
}
pub fn get_stats(&self) -> ModelPoolStats {
let (used, total) = self.get_memory_usage();
let loaded_models = self.model_slots.values().filter(|s| s.is_loaded).count();
let total_models = self.model_slots.len();
ModelPoolStats {
used_memory_mb: used / 1024 / 1024,
total_memory_mb: total / 1024 / 1024,
loaded_models,
total_models,
memory_usage_percent: self.get_memory_usage_percent(),
}
}
}
#[derive(Debug, Clone)]
pub struct ModelPoolStats {
pub used_memory_mb: u64,
pub total_memory_mb: u64,
pub loaded_models: usize,
pub total_models: usize,
pub memory_usage_percent: f64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_model_pool_creation() {
let config = ModelWeightPoolConfig::default();
let pool = ModelWeightPool::new("test_device".to_string(), config);
let (used, total) = pool.get_memory_usage();
assert_eq!(used, 0);
assert!(total > 0);
}
#[test]
fn test_allocate_model() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
let result = pool.allocate_for_model("model1", 512 * 1024 * 1024);
assert!(result.is_ok());
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 512 * 1024 * 1024);
}
#[test]
fn test_release_model() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
pool.release_model("model1");
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 0);
}
#[test]
fn test_can_load_model() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
cache_models: true,
..Default::default()
};
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
assert!(pool.can_load_model(512 * 1024 * 1024));
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
pool.mark_model_unloaded("model1");
assert!(pool.can_load_model(512 * 1024 * 1024));
}
#[test]
fn test_insufficient_memory() {
let config = ModelWeightPoolConfig {
max_memory_mb: 512,
cache_models: false,
..Default::default()
};
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
let result = pool.allocate_for_model("model2", 512 * 1024 * 1024);
assert!(result.is_err());
}
#[test]
fn test_update_usage() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
let slot_before = pool.get_model_slot("model1").unwrap();
let last_used_before = slot_before.last_used;
std::thread::sleep(std::time::Duration::from_millis(10));
pool.update_model_usage("model1");
let slot_after = pool.get_model_slot("model1").unwrap();
let last_used_after = slot_after.last_used;
assert!(last_used_after > last_used_before);
}
#[test]
fn test_stats() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = ModelWeightPool::new("test_device".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
pool.allocate_for_model("model2", 256 * 1024 * 1024)
.unwrap();
let stats = pool.get_stats();
assert_eq!(stats.used_memory_mb, 768);
assert_eq!(stats.total_memory_mb, 1024);
assert_eq!(stats.loaded_models, 2);
assert!((stats.memory_usage_percent - 75.0).abs() < 0.1);
}
#[test]
fn test_can_load_model_cache_disabled_insufficient() {
let config = ModelWeightPoolConfig {
max_memory_mb: 512,
cache_models: false,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
pool.mark_model_unloaded("model1");
assert!(!pool.can_load_model(1024 * 1024));
}
#[test]
fn test_can_load_model_cache_disabled_sufficient() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
cache_models: false,
..Default::default()
};
let pool = ModelWeightPool::new("test".to_string(), config);
assert!(pool.can_load_model(512 * 1024 * 1024));
assert!(!pool.can_load_model(2048 * 1024 * 1024));
}
#[test]
fn test_lru_reclaim_single_model() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
cache_models: true,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 512 * 1024 * 1024)
.unwrap();
pool.allocate_for_model("model2", 256 * 1024 * 1024)
.unwrap();
pool.mark_model_unloaded("model1");
let result = pool.allocate_for_model("model3", 512 * 1024 * 1024);
assert!(result.is_ok(), "LRU reclaim should succeed");
assert!(pool.get_model_slot("model1").is_none());
assert!(pool.get_model_slot("model2").is_some());
assert!(pool.get_model_slot("model3").is_some());
let stats = pool.get_stats();
assert_eq!(stats.total_models, 2);
assert_eq!(stats.loaded_models, 2);
}
#[test]
fn test_lru_reclaim_multiple_models() {
let config = ModelWeightPoolConfig {
max_memory_mb: 512,
cache_models: true,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 256 * 1024 * 1024)
.unwrap();
pool.allocate_for_model("model2", 256 * 1024 * 1024)
.unwrap();
pool.mark_model_unloaded("model1");
pool.mark_model_unloaded("model2");
let result = pool.allocate_for_model("model3", 512 * 1024 * 1024);
assert!(result.is_ok(), "multi-model reclaim should succeed");
assert!(pool.get_model_slot("model1").is_none());
assert!(pool.get_model_slot("model2").is_none());
assert!(pool.get_model_slot("model3").is_some());
}
#[test]
fn test_lru_reclaim_insufficient_returns_error() {
let config = ModelWeightPoolConfig {
max_memory_mb: 256,
cache_models: true,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 128 * 1024 * 1024)
.unwrap();
pool.allocate_for_model("model2", 128 * 1024 * 1024)
.unwrap();
pool.mark_model_unloaded("model1");
let result = pool.allocate_for_model("model3", 512 * 1024 * 1024);
assert!(result.is_err(), "should fail when reclaim insufficient");
}
#[test]
fn test_allocate_for_model_already_exists_noop() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 256 * 1024 * 1024)
.unwrap();
let (used_after_first, _) = pool.get_memory_usage();
assert_eq!(used_after_first, 256 * 1024 * 1024);
let result = pool.allocate_for_model("model1", 512 * 1024 * 1024);
assert!(result.is_ok());
let (used_after_second, _) = pool.get_memory_usage();
assert_eq!(
used_after_second, used_after_first,
"re-allocate should not change memory"
);
let slot = pool.get_model_slot("model1").unwrap();
assert_eq!(slot.memory_allocated, 256 * 1024 * 1024);
}
#[test]
fn test_release_model_nonexistent_noop() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test".to_string(), config);
let (used_before, _) = pool.get_memory_usage();
pool.release_model("nonexistent");
let (used_after, _) = pool.get_memory_usage();
assert_eq!(used_before, used_after);
}
#[test]
fn test_update_model_usage_nonexistent_noop() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.update_model_usage("nonexistent");
assert_eq!(pool.get_model_slots().len(), 0);
}
#[test]
fn test_mark_model_unloaded_nonexistent_noop() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.mark_model_unloaded("nonexistent");
assert_eq!(pool.get_model_slots().len(), 0);
}
#[test]
fn test_mark_model_unloaded_then_loaded_via_update() {
let config = ModelWeightPoolConfig::default();
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("model1", 128 * 1024 * 1024)
.unwrap();
assert!(pool.get_model_slot("model1").unwrap().is_loaded);
pool.mark_model_unloaded("model1");
assert!(!pool.get_model_slot("model1").unwrap().is_loaded);
assert_eq!(pool.get_stats().loaded_models, 0);
pool.update_model_usage("model1");
assert!(pool.get_model_slot("model1").unwrap().is_loaded);
assert_eq!(pool.get_stats().loaded_models, 1);
}
#[test]
fn test_get_model_slots_returns_all() {
let config = ModelWeightPoolConfig {
max_memory_mb: 2048,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("m1", 128 * 1024 * 1024).unwrap();
pool.allocate_for_model("m2", 256 * 1024 * 1024).unwrap();
pool.allocate_for_model("m3", 512 * 1024 * 1024).unwrap();
let slots = pool.get_model_slots();
assert_eq!(slots.len(), 3);
let names: Vec<String> = slots.iter().map(|s| s.model_name.clone()).collect();
assert!(names.contains(&"m1".to_string()));
assert!(names.contains(&"m2".to_string()));
assert!(names.contains(&"m3".to_string()));
}
#[test]
fn test_get_model_slot_nonexistent_returns_none() {
let config = ModelWeightPoolConfig::default();
let pool = ModelWeightPool::new("test".to_string(), config);
assert!(pool.get_model_slot("nonexistent").is_none());
}
#[test]
fn test_clear_resets_pool() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("m1", 256 * 1024 * 1024).unwrap();
pool.allocate_for_model("m2", 256 * 1024 * 1024).unwrap();
assert_eq!(pool.get_stats().total_models, 2);
assert_eq!(pool.get_memory_usage().0, 512 * 1024 * 1024);
pool.clear();
let stats = pool.get_stats();
assert_eq!(stats.total_models, 0);
assert_eq!(stats.loaded_models, 0);
assert_eq!(stats.used_memory_mb, 0);
let (used, _) = pool.get_memory_usage();
assert_eq!(used, 0);
}
#[test]
fn test_memory_usage_percent() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
assert!((pool.get_memory_usage_percent() - 0.0).abs() < 0.001);
pool.allocate_for_model("m1", 512 * 1024 * 1024).unwrap();
assert!((pool.get_memory_usage_percent() - 50.0).abs() < 0.001);
pool.allocate_for_model("m2", 512 * 1024 * 1024).unwrap();
assert!((pool.get_memory_usage_percent() - 100.0).abs() < 0.001);
}
#[test]
fn test_can_load_model_with_reclaimable_memory() {
let config = ModelWeightPoolConfig {
max_memory_mb: 1024,
cache_models: true,
..Default::default()
};
let mut pool = ModelWeightPool::new("test".to_string(), config);
pool.allocate_for_model("m1", 1024 * 1024 * 1024).unwrap();
pool.mark_model_unloaded("m1");
assert!(pool.can_load_model(512 * 1024 * 1024));
assert!(!pool.can_load_model(2048 * 1024 * 1024));
}
}