Aniruddha7/QueryLens-Text2SQL_DocVQA-V2
0
1"""2Memory management utilities for the text2sql agent.3"""4import gc5import os6import time7import psutil8import threading9import weakref10from contextlib import contextmanager11from typing import Optional, Callable, Dict, Any12 13class MemoryManager:14 """Memory manager for monitoring and managing system memory."""15 16 def __init__(self, threshold_mb: float = 1000.0, critical_mb: float = 300.0,17 check_interval: float = 0.5, monitor_active: bool = True):18 """Initialize the memory manager.19 20 Args:21 threshold_mb: Memory threshold in MB below which to take action22 critical_mb: Critical memory level in MB below which to abort operations23 check_interval: Interval in seconds for background memory checks24 monitor_active: Whether to start the background monitor thread25 """26 # Allow environment overrides to tune aggressiveness27 try:28 _env_threshold = os.environ.get("MEM_THRESHOLD_MB")29 _env_critical = os.environ.get("MEM_CRITICAL_MB")30 if _env_threshold: threshold_mb = float(_env_threshold)31 if _env_critical: critical_mb = float(_env_critical)32 except Exception:33 pass34 self.threshold_mb = threshold_mb35 self.critical_mb = critical_mb36 self.check_interval = check_interval37 self._lock = threading.RLock()38 self._monitor_thread = None39 self._keep_monitoring = False40 41 # Track instances with weak references when possible to avoid holding42 # strong references that prevent GC. Keep a fallback list for objects43 # that are not weak-referenceable.44 self._active_llm_refs = [] # list of weakref.ref where possible45 self._active_llm_strong = [] # fallback strong refs46 self._instance_count = 047 48 # Soft cap to avoid accidentally registering many parallel LLMs49 self.max_instances = 150 51 # Margin (MB) above critical required to allow new registrations52 self.registration_margin_mb = 200.053 54 # Track last emergency cleanup time to avoid repeated rapid-fire releases55 self._last_emergency = 0.056 self._emergency_cooldown = 10.0 # seconds between emergency attempts57 # Counter-based suspension flag to prevent emergency releases while in critical sections58 self._suspend_emergency_depth = 059 60 # Minimum available memory (MB) at which we will STILL retain a single LLM instance61 # even if below critical, to avoid thrashing (re-init cost is high). Configurable.62 try:63 self.min_retain_mb = float(os.environ.get("MIN_RETAIN_LLM_MB", "250"))64 except Exception:65 self.min_retain_mb = 250.066 67 # Start memory monitor if requested68 if monitor_active:69 self.start_monitor()70 71 def _prune_active_instances(self):72 """Remove dead weakrefs and return list of live instances."""73 live = []74 # Prune weak refs75 for ref in list(self._active_llm_refs):76 inst = ref()77 if inst is None:78 try:79 self._active_llm_refs.remove(ref)80 except ValueError:81 pass82 else:83 live.append(inst)84 # Add strong refs85 for inst in list(self._active_llm_strong):86 if inst is not None:87 live.append(inst)88 return live89 90 def register_llm_instance(self, llm_instance):91 """Register an LLM instance for tracking.92 93 This will refuse registration (and attempt to release the provided94 instance) if the system memory is too low. It prefers weak references95 so that the manager itself does not keep objects alive.96 """97 with self._lock:98 # Prune dead weakrefs first99 self._prune_active_instances()100 101 stats = self.get_memory_stats()102 # Relax policy: always allow the FIRST instance as long as we are above critical.103 # Only enforce margin for additional instances (even though max_instances is 1 today).104 if self._instance_count == 0:105 if stats['available_mb'] < self.critical_mb:106 # Too dangerous even for first instance107 try:108 if hasattr(llm_instance, 'close'):109 llm_instance.close()110 elif hasattr(llm_instance, 'release'):111 llm_instance.release()112 except Exception:113 pass114 raise MemoryError(f"Refusing first LLM registration: below critical ({stats['available_mb']:.1f}MB < {self.critical_mb}MB)")115 else:116 # For subsequent instances keep the stricter margin117 if stats['available_mb'] < (self.critical_mb + self.registration_margin_mb):118 try:119 if hasattr(llm_instance, 'close'):120 llm_instance.close()121 elif hasattr(llm_instance, 'release'):122 llm_instance.release()123 except Exception:124 pass125 raise MemoryError(f"Refusing to register additional LLM (avail={stats['available_mb']:.1f}MB; need >= {self.critical_mb + self.registration_margin_mb:.1f}MB)")126 127 # Enforce max instances cap128 if self._instance_count >= self.max_instances:129 raise MemoryError(f"Refusing to register LLM instance: max_instances={self.max_instances} reached")130 131 # Try to create a weakref to avoid holding a strong reference132 try:133 ref = weakref.ref(llm_instance, lambda r: None)134 self._active_llm_refs.append(ref)135 except TypeError:136 # Object not weak-referenceable; store strong ref but warn137 self._active_llm_strong.append(llm_instance)138 139 self._instance_count += 1140 print(f"LLM instance registered (total: {self._instance_count})")141 142 def unregister_llm_instance(self, llm_instance):143 """Unregister an LLM instance from tracking."""144 with self._lock:145 # Remove from weak refs146 removed = False147 for ref in list(self._active_llm_refs):148 inst = ref()149 if inst is None:150 # Dead ref; prune151 try:152 self._active_llm_refs.remove(ref)153 except ValueError:154 pass155 continue156 if inst is llm_instance:157 try:158 self._active_llm_refs.remove(ref)159 except ValueError:160 pass161 removed = True162 break163 164 # Remove from strong refs165 if not removed:166 for i, instance in enumerate(list(self._active_llm_strong)):167 if instance is llm_instance:168 del self._active_llm_strong[i]169 removed = True170 break171 172 if removed:173 self._instance_count = max(0, self._instance_count - 1)174 print(f"LLM instance unregistered (remaining: {self._instance_count})")175 176 def get_memory_stats(self) -> Dict[str, float]:177 """Get current memory statistics."""178 vm = psutil.virtual_memory()179 return {180 "total_mb": vm.total / (1024 * 1024),181 "available_mb": vm.available / (1024 * 1024),182 "used_mb": vm.used / (1024 * 1024),183 "percent_used": vm.percent184 }185 186 def is_memory_safe(self) -> bool:187 """Check if memory is at a safe level."""188 return self.get_memory_stats()["available_mb"] >= self.threshold_mb189 190 def is_memory_critical(self) -> bool:191 """Check if memory is at a critical level."""192 stats = self.get_memory_stats()193 available = stats["available_mb"]194 is_critical = available <= self.critical_mb195 196 # Debug logging for memory state197 if is_critical:198 print(f"🚨 [MEMORY_DEBUG] Critical state: {available:.1f}MB <= {self.critical_mb}MB threshold")199 200 return is_critical201 202 def force_collect_garbage(self):203 """Force aggressive garbage collection."""204 # Run multiple collection cycles205 for i in range(3):206 gc.collect(i)207 208 # Clear any reference cycles209 if hasattr(gc, 'freeze'):210 gc.freeze()211 212 # Clear module caches if possible213 if hasattr(gc, 'clear_caches'):214 gc.clear_caches()215 216 # Log memory status after cleanup217 stats = self.get_memory_stats()218 #print(f"After garbage collection: {stats['available_mb']:.2f}MB available ({stats['percent_used']}% used)")219 220 def _monitor_memory(self):221 """Background thread for memory monitoring."""222 while self._keep_monitoring:223 stats = self.get_memory_stats()224 225 # Check for critical memory conditions226 if stats["available_mb"] <= self.critical_mb:227 #print(f"⚠️ CRITICAL MEMORY WARNING: Only {stats['available_mb']:.2f}MB available!")228 self.force_collect_garbage()229 230 # If still critical after cleanup, take more drastic measures231 if self.is_memory_critical():232 now = time.time()233 # Only attempt emergency release if cooldown passed234 if now - self._last_emergency > self._emergency_cooldown:235 # Check if we have tracked instances (prune dead refs first)236 live = self._prune_active_instances()237 if live:238 # If only a single LLM instance and we still have above min_retain_mb, skip releasing to prevent constant re-init.239 if len(live) == 1 and stats["available_mb"] > self.min_retain_mb:240 print(f"⚠️ Critical memory but retaining single LLM (avail={stats['available_mb']:.1f}MB > {self.min_retain_mb}MB safeguard)")241 else:242 # Skip emergency release while suspended243 if self._suspend_emergency_depth > 0:244 print(f"🚫 Emergency release suppressed (depth={self._suspend_emergency_depth}) during critical section")245 else:246 self._emergency_release_llm_instances()247 else:248 # Log once per cooldown when no instances are registered249 print("🚨 Emergency memory state detected but no LLM instances registered to release")250 self._last_emergency = now251 else:252 # Cooldown active; skip repeated emergency attempts253 pass254 255 # Regular check for low memory256 elif stats["available_mb"] <= self.threshold_mb:257 #print(f"⚠️ Low memory warning: {stats['available_mb']:.2f}MB available")258 self.force_collect_garbage()259 260 # Sleep for the check interval261 time.sleep(self.check_interval)262 263 def _emergency_release_llm_instances(self):264 """Emergency release of LLM instances to free memory."""265 with self._lock:266 # Build list of live instances from weak and strong refs267 instances = self._prune_active_instances()268 print(f"Attempting to release {len(instances)} LLM instances")269 270 for llm in instances:271 try:272 # Try to close or release resources273 if hasattr(llm, 'close'):274 llm.close()275 elif hasattr(llm, 'release'):276 llm.release()277 278 # Set to None to break references279 if hasattr(llm, '_client'):280 llm._client = None281 if hasattr(llm, '_model'):282 llm._model = None283 except Exception as e:284 print(f"Error releasing LLM instance: {str(e)}")285 286 # Clear our tracking lists287 self._active_llm_refs.clear()288 self._active_llm_strong.clear()289 self._instance_count = 0290 291 # Force garbage collection after cleanup and attempt to return pages292 self.force_collect_garbage()293 try:294 # Best-effort working set trim to return pages to OS (Windows)295 self.trim_working_set()296 except Exception:297 pass298 299 # ---------------- Public high-level cleanup APIs -----------------300 def release_all_llms(self):301 """Public wrapper to release all tracked LLM instances (non-emergency manual call)."""302 self._emergency_release_llm_instances()303 304 def trim_working_set(self):305 """Attempt to return unused pages to OS (best-effort, platform-specific)."""306 try:307 import platform, ctypes308 if platform.system().lower() == 'windows':309 PROCESS_SET_QUOTA = 0x0100310 PROCESS_QUERY_INFORMATION = 0x0400311 kernel32 = ctypes.windll.kernel32312 psapi = ctypes.windll.psapi313 GetCurrentProcess = kernel32.GetCurrentProcess314 hProc = GetCurrentProcess()315 psapi.EmptyWorkingSet(hProc)316 else:317 # For *nix, reading /proc/self might encourage trimming after gc318 pass319 except Exception as e:320 print(f"Working set trim not supported/failed: {e}")321 322 def perform_full_cleanup(self) -> Dict[str, float]:323 """Comprehensive cleanup: release LLMs, GC, trim working set, return new stats."""324 before = self.get_memory_stats()325 self.release_all_llms()326 self.force_collect_garbage()327 self.trim_working_set()328 after = self.get_memory_stats()329 delta = after['available_mb'] - before['available_mb']330 print(f"Post-query full cleanup reclaimed {delta:.2f}MB (avail: {after['available_mb']:.2f}MB)")331 return after332 333 # --------------- Critical section helpers ---------------334 @contextmanager335 def suspend_emergency(self, reason: str = ""):336 """Context manager to temporarily suspend emergency LLM releases.337 338 Use around critical sections (e.g., active LLM generation) to prevent339 the monitor thread from tearing down the only LLM mid-call, which can340 cause intermittent generation failures.341 """342 with self._lock:343 self._suspend_emergency_depth += 1344 if reason:345 print(f"⏸️ Suspended emergency releases (depth={self._suspend_emergency_depth}) reason={reason}")346 try:347 yield348 finally:349 with self._lock:350 self._suspend_emergency_depth = max(0, self._suspend_emergency_depth - 1)351 if reason:352 print(f"▶️ Resumed emergency releases (depth={self._suspend_emergency_depth}) reason={reason}")353 354 def start_monitor(self):355 """Start the background memory monitor."""356 with self._lock:357 if self._monitor_thread is None or not self._monitor_thread.is_alive():358 self._keep_monitoring = True359 self._monitor_thread = threading.Thread(360 target=self._monitor_memory,361 daemon=True,362 name="MemoryMonitorThread"363 )364 self._monitor_thread.start()365 print("Memory monitor started in background")366 367 def stop_monitor(self):368 """Stop the background memory monitor."""369 with self._lock:370 self._keep_monitoring = False371 if self._monitor_thread and self._monitor_thread.is_alive():372 self._monitor_thread.join(timeout=1.0)373 print("Memory monitor stopped")374 375# Create a global memory manager instance376memory_manager = MemoryManager(377 threshold_mb=500.0, # 500MB — warning/action threshold (reduced from 1000MB)378 critical_mb=300.0, # 300MB — BSOD-protection hard floor (unchanged)379 check_interval=1.0, # Check every second380 monitor_active=True # Start monitoring immediately381)382 383class TimeoutManager:384 """Manages operation timeouts with proper cleanup."""385 386 @staticmethod387 def run_with_timeout(func: Callable, timeout: float, *args, **kwargs) -> Any:388 """389 Run a function with a timeout.390 391 Args:392 func: Function to run393 timeout: Timeout in seconds394 *args, **kwargs: Arguments to pass to the function395 396 Returns:397 The result of the function398 399 Raises:400 TimeoutError: If the function times out401 """402 result = None403 exception = None404 completed = threading.Event()405 406 # Thread function407 def worker():408 nonlocal result, exception409 try:410 result = func(*args, **kwargs)411 except Exception as e:412 exception = e413 finally:414 completed.set()415 416 # Start the worker thread417 thread = threading.Thread(target=worker, daemon=True)418 thread.start()419 420 # Wait for completion or timeout421 if not completed.wait(timeout):422 # Force garbage collection before raising timeout423 memory_manager.force_collect_garbage()424 425 # Raise timeout error426 raise TimeoutError(f"Operation timed out after {timeout} seconds")427 428 # If there was an exception, raise it429 if exception is not None:430 raise exception431 432 return result433 434def safe_check_memory(threshold_mb: Optional[float] = None) -> bool:435 """436 Safely check if there's enough memory available.437 Also performs garbage collection if memory is low.438 439 Args:440 threshold_mb: Optional threshold in MB (uses memory manager's threshold if None)441 442 Returns:443 True if memory is safe, False if it's below threshold444 """445 # Use memory manager's threshold if none provided446 if threshold_mb is None:447 threshold_mb = memory_manager.threshold_mb448 449 # Get current memory stats450 stats = memory_manager.get_memory_stats()451 452 # Log memory status453 print(f"Memory status: {stats['available_mb']:.2f}MB available of {stats['total_mb']:.2f}MB total")454 455 # If memory is below threshold, run garbage collection456 if stats['available_mb'] < threshold_mb:457 memory_manager.force_collect_garbage()458 459 # Get updated stats after garbage collection460 stats = memory_manager.get_memory_stats()461 462 # If still below threshold, return False463 if stats['available_mb'] < threshold_mb:464 print(f"⚠️ Low memory warning: Only {stats['available_mb']:.2f}MB available, need at least {threshold_mb:.2f}MB")465 return False466 467 return True468 