-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel_cache.py
More file actions
44 lines (38 loc) · 1.68 KB
/
Copy pathmodel_cache.py
File metadata and controls
44 lines (38 loc) · 1.68 KB
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
# 创建一个全局模型管理器来缓存已加载的模型, 避免重复创建
class ModelCache:
def __init__(self):
self._detector = None
self._detector_device = None
self._embedder = None
self._embedder_device = None
self._embedder_source = None
self._databases = {} # 缓存已加载的数据库
def get_detector(self, device: str):
if self._detector is None or self._detector_device != device:
from modules.face_detection import load_detector
self._detector = load_detector(device)
self._detector_device = device
return self._detector
def get_embedder(self, device: str, pretrained_source: str):
if (self._embedder is None or
self._embedder_device != device or
self._embedder_source != pretrained_source):
from modules.feature_extraction import load_embedder
self._embedder = load_embedder(device, pretrained_source)
self._embedder_device = device
self._embedder_source = pretrained_source
return self._embedder
def get_database(self, path: str):
"""获取缓存的数据库实例,如果不存在则加载"""
if path not in self._databases:
from modules.matcher import load_db
self._databases[path] = load_db(path)
return self._databases[path]
def invalidate_database(self, path: str):
"""使指定路径的数据库缓存失效"""
if path in self._databases:
del self._databases[path]
def clear_database_cache(self):
"""清空所有数据库缓存"""
self._databases.clear()
model_cache = ModelCache()