File size: 3,193 Bytes
b611f38
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
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)