import gc import logging import threading from typing import Dict, Optional, List, Any, TYPE_CHECKING if TYPE_CHECKING: from adapters.base import BaseOCRAdapter logger = logging.getLogger("OCRMemoryManager") class OCRModelManager: """ Manages lazy loading and unloading of OCR adapters. Keeps at most `max_cached_models` in memory to prevent out-of-memory (OOM) errors. """ def __init__(self, max_cached_models: int = 1): self.max_cached_models = max_cached_models self.adapters: Dict[str, Any] = {} self.loaded_order: List[str] = [] self._lock = threading.Lock() def register_adapter(self, key: str, adapter: Any) -> None: """Register an adapter instance in the manager.""" self.adapters[key] = adapter def get_adapter(self, key: str) -> Any: """Retrieve an adapter by key.""" if key not in self.adapters: raise KeyError(f"Adapter '{key}' is not registered. Available: {list(self.adapters.keys())}") return self.adapters[key] def load_and_acquire(self, key: str) -> Any: """ Loads the requested adapter, evicting the least-recently used adapter if the cache limit is reached. """ with self._lock: if key not in self.adapters: raise KeyError(f"Adapter '{key}' not found.") adapter = self.adapters[key] # If already loaded, move it to the end of loaded_order if adapter.is_loaded: if key in self.loaded_order: self.loaded_order.remove(key) self.loaded_order.append(key) return adapter # Evict models if capacity reached while len(self.loaded_order) >= self.max_cached_models: evict_key = self.loaded_order.pop(0) logger.info(f"Evicting model '{evict_key}' from memory to make room for '{key}'") try: self.adapters[evict_key].unload_model() except Exception as e: logger.warning(f"Error while unloading '{evict_key}': {e}") # Load the requested adapter logger.info(f"Loading model '{key}' into memory/VRAM...") adapter.load_model() adapter._is_loaded = True self.loaded_order.append(key) return adapter def unload_all(self) -> None: """Unloads all loaded adapters and flushes memory.""" with self._lock: for key, adapter in self.adapters.items(): if adapter.is_loaded: try: adapter.unload_model() except Exception as e: logger.warning(f"Error unloading '{key}': {e}") self.loaded_order.clear() gc.collect() try: import torch if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.ipc_collect() except ImportError: pass # Global singleton instance model_manager = OCRModelManager(max_cached_models=1)