import os
|
|
from functools import lru_cache
|
|
from typing import Iterator, Optional, Tuple
|
|
|
|
from ultralytics import YOLO
|
|
|
|
|
|
MODEL_MAP = {
|
|
"floating": "model/Floating/best.pt",
|
|
"gate": "model/gate_model/best.pt",
|
|
"shore_garbage": "model/shore_garbage_model/best.pt",
|
|
"water_gauge": "model/water_gauge_model/best.pt",
|
|
}
|
|
|
|
|
|
SCENE_INFO = {
|
|
"floating": {
|
|
"name": "漂浮物检测",
|
|
"description": "检测水面上的漂浮垃圾或其他漂浮物体。",
|
|
},
|
|
"gate": {
|
|
"name": "闸口场景检测",
|
|
"description": "检测闸口区域周边的目标和异常情况。",
|
|
},
|
|
"shore_garbage": {
|
|
"name": "岸边垃圾检测",
|
|
"description": "检测岸线附近堆积或散落的垃圾。",
|
|
},
|
|
"water_gauge": {
|
|
"name": "水尺检测",
|
|
"description": "检测水位监测场景中的水尺目标。",
|
|
},
|
|
}
|
|
|
|
|
|
def get_model_path(model_type: str) -> str:
|
|
return os.path.join(os.path.dirname(__file__), MODEL_MAP[model_type])
|
|
|
|
|
|
@lru_cache(maxsize=len(MODEL_MAP))
|
|
def load_model(model_type: str) -> YOLO:
|
|
return YOLO(get_model_path(model_type))
|
|
|
|
|
|
def iter_models(model_type: Optional[str] = None) -> Iterator[Tuple[str, YOLO]]:
|
|
if model_type is not None:
|
|
yield model_type, load_model(model_type)
|
|
return
|
|
|
|
for current_model_type in MODEL_MAP:
|
|
yield current_model_type, load_model(current_model_type)
|