Files
foxhunt/ml/src/safety/memory_manager.rs
2026-02-23 01:03:44 +01:00

595 lines
20 KiB
Rust

//! Safe Memory Management for ML Operations
//!
//! This module provides comprehensive memory management and monitoring
//! to prevent OOM conditions and memory leaks in ML operations.
#![deny(clippy::unwrap_used)]
#![deny(clippy::expect_used)]
#![deny(clippy::panic)]
use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Instant;
use candle_core::Device;
use tracing::{debug, error, info, warn};
use super::{MLSafetyConfig, MLSafetyError, SafetyResult, SafetyStatus};
/// Memory usage tracking per device
#[derive(Debug)]
struct DeviceMemoryUsage {
allocated_bytes: AtomicUsize,
peak_bytes: AtomicUsize,
allocation_count: AtomicUsize,
last_cleanup: Instant,
}
impl DeviceMemoryUsage {
fn new() -> Self {
Self {
allocated_bytes: AtomicUsize::new(0),
peak_bytes: AtomicUsize::new(0),
allocation_count: AtomicUsize::new(0),
last_cleanup: Instant::now(),
}
}
fn allocate(&self, bytes: usize) -> usize {
let new_total = self.allocated_bytes.fetch_add(bytes, Ordering::Relaxed) + bytes;
self.allocation_count.fetch_add(1, Ordering::Relaxed);
// Update peak if necessary
let current_peak = self.peak_bytes.load(Ordering::Relaxed);
if new_total > current_peak {
self.peak_bytes.store(new_total, Ordering::Relaxed);
}
new_total
}
fn deallocate(&self, bytes: usize) -> usize {
self.allocated_bytes.fetch_sub(
bytes.min(self.allocated_bytes.load(Ordering::Relaxed)),
Ordering::Relaxed,
)
}
fn get_allocated(&self) -> usize {
self.allocated_bytes.load(Ordering::Relaxed)
}
fn get_peak(&self) -> usize {
self.peak_bytes.load(Ordering::Relaxed)
}
fn get_allocation_count(&self) -> usize {
self.allocation_count.load(Ordering::Relaxed)
}
fn reset_peak(&self) {
let current = self.allocated_bytes.load(Ordering::Relaxed);
self.peak_bytes.store(current, Ordering::Relaxed);
}
}
/// Safe memory manager with comprehensive monitoring
pub struct SafeMemoryManager {
config: MLSafetyConfig,
device_usage: HashMap<String, DeviceMemoryUsage>,
system_memory_limit: usize,
cleanup_threshold: f64,
emergency_cleanup_callbacks: Vec<Box<dyn Fn() + Send + Sync>>,
}
impl std::fmt::Debug for SafeMemoryManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SafeMemoryManager")
.field("config", &self.config)
.field("device_usage", &self.device_usage)
.field("system_memory_limit", &self.system_memory_limit)
.field("cleanup_threshold", &self.cleanup_threshold)
.field(
"emergency_cleanup_callbacks",
&format!("{} callbacks", self.emergency_cleanup_callbacks.len()),
)
.finish()
}
}
impl SafeMemoryManager {
/// Create new safe memory manager
pub fn new(config: &MLSafetyConfig) -> Self {
Self {
config: config.clone(),
device_usage: HashMap::new(),
system_memory_limit: 32 * 1024 * 1024 * 1024, // 32GB default
cleanup_threshold: 0.85, // 85% usage triggers cleanup
emergency_cleanup_callbacks: Vec::new(),
}
}
/// Check memory availability before allocation
pub fn check_memory_availability(
&mut self,
requested_bytes: usize,
device: &Device,
) -> SafetyResult<()> {
let device_key = self.device_key(device);
// Get or create device usage tracker
let usage = self
.device_usage
.entry(device_key.clone())
.or_insert_with(DeviceMemoryUsage::new);
let current_usage = usage.get_allocated();
let projected_usage = current_usage + requested_bytes;
// Check device-specific limits
match device {
Device::Cpu => {
if projected_usage > self.system_memory_limit {
return Err(MLSafetyError::MemorySafety {
reason: format!(
"CPU memory limit exceeded: {} + {} = {} > {} limit",
self.format_bytes(current_usage),
self.format_bytes(requested_bytes),
self.format_bytes(projected_usage),
self.format_bytes(self.system_memory_limit)
),
});
}
},
Device::Cuda(_) => {
if projected_usage > self.config.max_gpu_memory_bytes {
return Err(MLSafetyError::MemorySafety {
reason: format!(
"GPU memory limit exceeded: {} + {} = {} > {} limit",
self.format_bytes(current_usage),
self.format_bytes(requested_bytes),
self.format_bytes(projected_usage),
self.format_bytes(self.config.max_gpu_memory_bytes)
),
});
}
},
Device::Metal(_) => {
// Metal device memory checking
if projected_usage > self.config.max_gpu_memory_bytes {
return Err(MLSafetyError::MemorySafety {
reason: format!(
"Metal memory limit exceeded: {} + {} = {} > {} limit",
self.format_bytes(current_usage),
self.format_bytes(requested_bytes),
self.format_bytes(projected_usage),
self.format_bytes(self.config.max_gpu_memory_bytes)
),
});
}
},
}
// Check if cleanup is needed
let usage_ratio = projected_usage as f64 / self.get_memory_limit(device) as f64;
if usage_ratio > self.cleanup_threshold {
warn!(
"Memory usage high on {}: {:.1}% (threshold: {:.1}%)",
device_key,
usage_ratio * 100.0,
self.cleanup_threshold * 100.0
);
// Trigger automatic cleanup if enabled
if self.config.auto_fallback {
warn!(
"Memory usage high, cleanup needed for device: {}",
device_key
);
// Note: Cleanup would be triggered asynchronously in real implementation
for callback in &self.emergency_cleanup_callbacks {
callback();
}
}
}
debug!(
"Memory check passed for {}: {} available, {} requested",
device_key,
self.format_bytes(self.get_memory_limit(device) - current_usage),
self.format_bytes(requested_bytes)
);
Ok(())
}
/// Record memory allocation
pub fn record_allocation(&mut self, bytes: usize, device: &Device) -> usize {
let device_key = self.device_key(device);
let usage = self
.device_usage
.entry(device_key.clone())
.or_insert_with(DeviceMemoryUsage::new);
let new_total = usage.allocate(bytes);
debug!(
"Memory allocated on {}: {} bytes, total: {}",
device_key,
self.format_bytes(bytes),
self.format_bytes(new_total)
);
new_total
}
/// Record memory deallocation
pub fn record_deallocation(&mut self, bytes: usize, device: &Device) -> usize {
let device_key = self.device_key(device);
if let Some(usage) = self.device_usage.get(&device_key) {
let new_total = usage.deallocate(bytes);
debug!(
"Memory deallocated on {}: {} bytes, remaining: {}",
device_key,
self.format_bytes(bytes),
self.format_bytes(new_total)
);
new_total
} else {
warn!(
"Attempted to deallocate from untracked device: {}",
device_key
);
0
}
}
/// Get current memory usage for device
pub fn get_memory_usage(&self, device: &Device) -> usize {
let device_key = self.device_key(device);
self.device_usage
.get(&device_key)
.map(|usage| usage.get_allocated())
.unwrap_or(0)
}
/// Get peak memory usage for device
pub fn get_peak_memory_usage(&self, device: &Device) -> usize {
let device_key = self.device_key(device);
self.device_usage
.get(&device_key)
.map(|usage| usage.get_peak())
.unwrap_or(0)
}
/// Get memory usage statistics
pub fn get_memory_stats(&self) -> HashMap<String, HashMap<String, String>> {
let mut stats = HashMap::new();
for (device_key, usage) in &self.device_usage {
let mut device_stats = HashMap::new();
device_stats.insert(
"allocated".to_string(),
self.format_bytes(usage.get_allocated()),
);
device_stats.insert("peak".to_string(), self.format_bytes(usage.get_peak()));
device_stats.insert(
"allocation_count".to_string(),
usage.get_allocation_count().to_string(),
);
let limit = if device_key.contains("cpu") {
self.system_memory_limit
} else {
self.config.max_gpu_memory_bytes
};
device_stats.insert("limit".to_string(), self.format_bytes(limit));
let usage_percent = (usage.get_allocated() as f64 / limit as f64) * 100.0;
device_stats.insert(
"usage_percent".to_string(),
format!("{:.1}%", usage_percent),
);
stats.insert(device_key.clone(), device_stats);
}
stats
}
/// Check overall memory safety status
pub async fn get_status(&self) -> SafetyStatus {
let mut warnings = Vec::new();
let mut dangers = Vec::new();
for (device_key, usage) in &self.device_usage {
let limit = if device_key.contains("cpu") {
self.system_memory_limit
} else {
self.config.max_gpu_memory_bytes
};
let usage_ratio = usage.get_allocated() as f64 / limit as f64;
if usage_ratio >= 0.95 {
dangers.push(format!(
"{}: {:.1}% usage (critical)",
device_key,
usage_ratio * 100.0
));
} else if usage_ratio > self.cleanup_threshold {
warnings.push(format!(
"{}: {:.1}% usage (high)",
device_key,
usage_ratio * 100.0
));
}
}
if !dangers.is_empty() {
SafetyStatus::Critical {
reason: format!("Critical memory usage: {}", dangers.join(", ")),
}
} else if !warnings.is_empty() {
SafetyStatus::Warning {
reason: format!("High memory usage: {}", warnings.join(", ")),
}
} else {
SafetyStatus::Safe
}
}
/// Trigger memory cleanup for a device
async fn trigger_cleanup(&mut self, device_key: &str) -> SafetyResult<()> {
info!("Triggering memory cleanup for device: {}", device_key);
// Execute cleanup callbacks
for callback in &self.emergency_cleanup_callbacks {
callback();
}
// Reset peak tracking
if let Some(usage) = self.device_usage.get(device_key) {
usage.reset_peak();
}
// No garbage collector needed: Rust's ownership model and RAII handle
// deallocation automatically when Tensors and VarMaps go out of scope.
// GPU memory (CUDA/Metal) is freed when candle Tensors are dropped.
info!("Memory cleanup completed for device: {}", device_key);
Ok(())
}
/// Emergency cleanup - clear all tracked memory
pub async fn emergency_cleanup(&mut self) -> SafetyResult<()> {
error!("Emergency memory cleanup initiated");
// Execute all cleanup callbacks
for callback in &self.emergency_cleanup_callbacks {
callback();
}
// Reset all memory tracking
for (device_key, usage) in &self.device_usage {
let allocated = usage.get_allocated();
if allocated > 0 {
warn!(
"Emergency cleanup: {} had {} allocated",
device_key,
self.format_bytes(allocated)
);
}
usage.allocated_bytes.store(0, Ordering::Relaxed);
usage.reset_peak();
}
info!("Emergency memory cleanup completed");
Ok(())
}
/// Add emergency cleanup callback
pub fn add_cleanup_callback<F>(&mut self, callback: F)
where
F: Fn() + Send + Sync + 'static,
{
self.emergency_cleanup_callbacks.push(Box::new(callback));
}
/// Set system memory limit
pub fn set_system_memory_limit(&mut self, bytes: usize) {
self.system_memory_limit = bytes;
info!("System memory limit set to: {}", self.format_bytes(bytes));
}
/// Set cleanup threshold (0.0 to 1.0)
pub fn set_cleanup_threshold(&mut self, threshold: f64) {
self.cleanup_threshold = threshold.clamp(0.0, 1.0);
info!("Cleanup threshold set to: {:.1}%", threshold * 100.0);
}
/// Get device-specific memory limit
fn get_memory_limit(&self, device: &Device) -> usize {
match device {
Device::Cpu => self.system_memory_limit,
Device::Cuda(_) | Device::Metal(_) => self.config.max_gpu_memory_bytes,
}
}
/// Generate device key for tracking
fn device_key(&self, device: &Device) -> String {
match device {
Device::Cpu => "cpu".to_string(),
Device::Cuda(id) => format!("cuda_{:?}", id),
Device::Metal(id) => format!("metal_{:?}", id),
}
}
/// Format bytes in human-readable form
fn format_bytes(&self, bytes: usize) -> String {
const UNITS: &[&str] = &["B", "KB", "MB", "GB", "TB"];
const THRESHOLD: f64 = 1024.0;
if bytes == 0 {
return "0 B".to_string();
}
let mut size = bytes as f64;
let mut unit_index = 0;
while size >= THRESHOLD && unit_index < UNITS.len() - 1 {
size /= THRESHOLD;
unit_index += 1;
}
if unit_index == 0 {
format!("{} {}", bytes, UNITS.get(unit_index).unwrap_or(&"B"))
} else {
format!("{:.1} {}", size, UNITS.get(unit_index).unwrap_or(&"B"))
}
}
/// Reset memory tracking for device
pub fn reset_device_tracking(&mut self, device: &Device) {
let device_key = self.device_key(device);
self.device_usage.remove(&device_key);
debug!("Reset memory tracking for device: {}", device_key);
}
/// Reset all memory tracking
pub fn reset_all_tracking(&mut self) {
self.device_usage.clear();
debug!("Reset all memory tracking");
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
fn create_test_manager() -> SafeMemoryManager {
SafeMemoryManager::new(&MLSafetyConfig::default())
}
#[test]
fn test_memory_allocation_tracking() {
let mut manager = create_test_manager();
let device = Device::Cpu;
// Record allocation
let total = manager.record_allocation(1024, &device);
assert_eq!(total, 1024);
assert_eq!(manager.get_memory_usage(&device), 1024);
// Record more allocation
manager.record_allocation(512, &device);
assert_eq!(manager.get_memory_usage(&device), 1536);
// Record deallocation
manager.record_deallocation(512, &device);
assert_eq!(manager.get_memory_usage(&device), 1024);
}
#[test]
fn test_memory_limit_checking() {
let mut manager = create_test_manager();
manager.set_system_memory_limit(2048); // 2KB limit for testing
let device = Device::Cpu;
// Should pass - under limit
assert!(manager.check_memory_availability(1024, &device).is_ok());
// Should fail - over limit
assert!(manager.check_memory_availability(3072, &device).is_err());
}
#[test]
fn test_peak_tracking() {
let mut manager = create_test_manager();
let device = Device::Cpu;
// Allocate and check peak
manager.record_allocation(1024, &device);
assert_eq!(manager.get_peak_memory_usage(&device), 1024);
// Allocate more and check peak updates
manager.record_allocation(512, &device);
assert_eq!(manager.get_peak_memory_usage(&device), 1536);
// Deallocate and check peak remains
manager.record_deallocation(512, &device);
assert_eq!(manager.get_peak_memory_usage(&device), 1536);
assert_eq!(manager.get_memory_usage(&device), 1024);
}
#[test]
fn test_byte_formatting() {
let manager = create_test_manager();
assert_eq!(manager.format_bytes(0), "0 B");
assert_eq!(manager.format_bytes(512), "512 B");
assert_eq!(manager.format_bytes(1024), "1.0 KB");
assert_eq!(manager.format_bytes(1536), "1.5 KB");
assert_eq!(manager.format_bytes(1024 * 1024), "1.0 MB");
assert_eq!(manager.format_bytes(1024 * 1024 * 1024), "1.0 GB");
}
#[test]
fn test_device_keys() {
let manager = create_test_manager();
assert_eq!(manager.device_key(&Device::Cpu), "cpu");
// Note: CUDA and Metal device testing requires actual device creation
// which is platform-specific and may not be available in all test environments.
// The device_key method uses Debug formatting which works for all device types.
}
#[tokio::test]
async fn test_cleanup_callback() {
let mut manager = create_test_manager();
let cleanup_called = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let cleanup_called_clone = cleanup_called.clone();
manager.add_cleanup_callback(move || {
cleanup_called_clone.store(true, Ordering::Relaxed);
});
// Trigger emergency cleanup
let cleanup_result = manager.emergency_cleanup().await;
assert!(cleanup_result.is_ok());
assert!(cleanup_called.load(Ordering::Relaxed));
}
#[tokio::test]
async fn test_safety_status() {
let mut manager = create_test_manager();
manager.set_system_memory_limit(1000); // Small limit for testing
let device = Device::Cpu;
// Safe status with low usage
manager.record_allocation(100, &device);
let status = manager.get_status().await;
assert!(matches!(status, SafetyStatus::Safe));
// Warning status with high usage
manager.record_allocation(800, &device); // 90% usage
let status = manager.get_status().await;
assert!(matches!(status, SafetyStatus::Warning { .. }));
// Critical status with very high usage
manager.record_allocation(50, &device); // 95% usage
let status = manager.get_status().await;
assert!(matches!(status, SafetyStatus::Critical { .. }));
}
}