OCR-Demo / utils /memory_manager.py
thangvckeygen's picture
Deploy 7-Model OCR Benchmark Space with Bounding Box Visualizer
b611f38
Raw
History Blame Contribute Delete
3.19 kB
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)