Spaces:
Running
Running
| 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) | |