Team Ai
Apppublic

Aniruddha7/QueryLens-Text2SQL_DocVQA-V2

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
memory_manager.py468 linesDownload Raw Back to Agent
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