595 lines
20 KiB
Rust
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 { .. }));
|
|
}
|
|
}
|