File size: 2,221 Bytes
0e757f3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# modelManager.py - Model Cache Manager
# PRODUCTION HARDENED

import time
import gc
import psutil
from typing import Dict, Optional, Any

class ModelManager:
    """Memory-aware model management with LRU cache."""
    
    def __init__(self, max_models: int = 3, memory_limit_mb: int = 6000):
        self.models: Dict[str, Any] = {}
        self.load_times: Dict[str, float] = {}
        self.use_counts: Dict[str, int] = {}
        self.max_models = max_models
        self.memory_limit_mb = memory_limit_mb
    
    def get_model(self, name: str, loader_func) -> Optional[Any]:
        """Get or load a model with caching."""
        # Check if already loaded
        if name in self.models:
            self.use_counts[name] = self.use_counts.get(name, 0) + 1
            return self.models[name]
        
        # Check memory before loading
        try:
            mem = psutil.virtual_memory()
            if mem.percent > 85:
                self._evict_models()
        except:
            pass
        
        # Load model
        start = time.time()
        try:
            session = loader_func(name)
            load_time = time.time() - start
            self.models[name] = session
            self.load_times[name] = load_time
            self.use_counts[name] = 1
            
            # Evict if over limit
            if len(self.models) > self.max_models:
                self._evict_models()
            
            return session
        except Exception as e:
            return None
    
    def _evict_models(self):
        """Evict least recently used models."""
        if len(self.models) <= 1:
            return
        
        sorted_models = sorted(
            self.models.keys(),
            key=lambda x: self.use_counts.get(x, 0)
        )
        
        to_remove = sorted_models[0]
        del self.models[to_remove]
        if to_remove in self.use_counts:
            del self.use_counts[to_remove]
        gc.collect()
    
    def get_stats(self) -> Dict:
        return {
            'loaded_models': list(self.models.keys()),
            'model_count': len(self.models),
            'use_counts': self.use_counts,
            'load_times': self.load_times
        }