26-7-31-1
This commit is contained in:
@@ -0,0 +1,246 @@
|
||||
"""YOLO 检测:管理员配置的权重优先;无权重或加载失败则演示模式。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import random
|
||||
from typing import Any
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from app.config import get_settings
|
||||
from app.schemas.models import Detection
|
||||
from app.services.yolo_manager import get_yolo_manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_YOLO_OK = False
|
||||
try:
|
||||
from ultralytics import YOLO # type: ignore
|
||||
|
||||
_YOLO_OK = True
|
||||
except Exception: # pragma: no cover
|
||||
YOLO = None # type: ignore
|
||||
logger.info("ultralytics 未安装,影像检测将使用演示模式")
|
||||
|
||||
|
||||
LABEL_ZH = {
|
||||
# 演示 / 规则降级
|
||||
"opacity": "片状阴影/渗出",
|
||||
"nodule": "结节",
|
||||
"fracture": "骨折线",
|
||||
"effusion": "积液",
|
||||
"lesion": "异常 dens 区",
|
||||
"calcification": "钙化",
|
||||
"mass": "占位",
|
||||
"object": "可疑区域",
|
||||
# 肺炎单类权重
|
||||
"Pneumonia": "肺炎",
|
||||
"pneumonia": "肺炎",
|
||||
# 胸部 X 光多病灶检测(VinBigData / 类似类别)
|
||||
"Aortic enlargement": "主动脉增宽",
|
||||
"Atelectasis": "肺不张",
|
||||
"Calcification": "钙化",
|
||||
"Cardiomegaly": "心脏增大",
|
||||
"Consolidation": "实变",
|
||||
"ILD": "间质性肺病",
|
||||
"Infiltration": "浸润",
|
||||
"Lung Opacity": "肺野透过度减低",
|
||||
"Nodule/Mass": "结节/肿块",
|
||||
"Other lesion": "其他病灶",
|
||||
"Pleural effusion": "胸腔积液",
|
||||
"Pleural thickening": "胸膜增厚",
|
||||
"Pneumothorax": "气胸",
|
||||
"Pulmonary fibrosis": "肺纤维化",
|
||||
}
|
||||
|
||||
|
||||
def yolo_available() -> bool:
|
||||
return _YOLO_OK
|
||||
|
||||
|
||||
class YoloDetector:
|
||||
def __init__(self) -> None:
|
||||
self._model = None
|
||||
self._mode: str = "demo"
|
||||
self._loaded_path: str | None = None
|
||||
self._class_names: dict[int, str] = {}
|
||||
self._load_error: str | None = None
|
||||
self.reload()
|
||||
|
||||
def reload(self) -> dict[str, Any]:
|
||||
"""按管理端配置重新加载权重。"""
|
||||
mgr = get_yolo_manager()
|
||||
mode_pref = mgr.effective_demo_mode()
|
||||
weights = mgr.active_weight_path()
|
||||
self._model = None
|
||||
self._class_names = {}
|
||||
self._load_error = None
|
||||
self._loaded_path = None
|
||||
|
||||
if mode_pref == "demo":
|
||||
self._mode = "demo"
|
||||
logger.info("YOLO 强制演示模式")
|
||||
return self.info()
|
||||
|
||||
if weights is None:
|
||||
self._mode = "demo"
|
||||
if mode_pref == "real":
|
||||
self._load_error = "已选 real 模式但未配置有效权重文件"
|
||||
logger.warning(self._load_error)
|
||||
else:
|
||||
logger.info("未配置权重,使用演示模式")
|
||||
return self.info()
|
||||
|
||||
if not _YOLO_OK:
|
||||
self._mode = "demo"
|
||||
self._load_error = "未安装 ultralytics,无法加载真实权重"
|
||||
logger.warning(self._load_error)
|
||||
return self.info()
|
||||
|
||||
try:
|
||||
self._model = YOLO(str(weights))
|
||||
self._mode = "real"
|
||||
self._loaded_path = str(weights)
|
||||
names = getattr(self._model, "names", None) or {}
|
||||
if isinstance(names, dict):
|
||||
self._class_names = {int(k): str(v) for k, v in names.items()}
|
||||
logger.info("已加载 YOLO 权重: %s", weights)
|
||||
except Exception as e:
|
||||
self._mode = "demo"
|
||||
self._model = None
|
||||
self._load_error = f"加载权重失败: {e}"
|
||||
logger.warning(self._load_error)
|
||||
return self.info()
|
||||
|
||||
@property
|
||||
def mode(self) -> str:
|
||||
return self._mode
|
||||
|
||||
def info(self) -> dict[str, Any]:
|
||||
return {
|
||||
"runtime_mode": self._mode,
|
||||
"ultralytics_installed": _YOLO_OK,
|
||||
"loaded_path": self._loaded_path,
|
||||
"class_names": list(self._class_names.values()) if self._class_names else [],
|
||||
"class_count": len(self._class_names),
|
||||
"load_error": self._load_error,
|
||||
"env_demo_mode": get_settings().demo_mode,
|
||||
}
|
||||
|
||||
def detect(self, image_bgr: np.ndarray, study_type: str = "CT") -> list[Detection]:
|
||||
if self._mode == "real" and self._model is not None:
|
||||
return self._detect_real(image_bgr)
|
||||
return self._detect_demo(image_bgr, study_type)
|
||||
|
||||
def _detect_real(self, image_bgr: np.ndarray) -> list[Detection]:
|
||||
results = self._model.predict(source=image_bgr, verbose=False)
|
||||
detections: list[Detection] = []
|
||||
if not results:
|
||||
return detections
|
||||
r0 = results[0]
|
||||
names = r0.names or self._class_names or {}
|
||||
boxes = getattr(r0, "boxes", None)
|
||||
if boxes is None:
|
||||
return detections
|
||||
for box in boxes:
|
||||
xyxy = box.xyxy[0].tolist()
|
||||
conf = float(box.conf[0]) if box.conf is not None else 0.0
|
||||
cls_id = int(box.cls[0]) if box.cls is not None else 0
|
||||
label = str(names.get(cls_id, f"class_{cls_id}"))
|
||||
detections.append(
|
||||
Detection(
|
||||
label=label,
|
||||
label_zh=LABEL_ZH.get(label, label),
|
||||
confidence=round(conf, 4),
|
||||
bbox=[round(float(x), 2) for x in xyxy],
|
||||
)
|
||||
)
|
||||
return detections
|
||||
|
||||
def _detect_demo(self, image_bgr: np.ndarray, study_type: str) -> list[Detection]:
|
||||
h, w = image_bgr.shape[:2]
|
||||
rng = random.Random(h * 31 + w * 17 + hash(study_type) % 997)
|
||||
catalog = {
|
||||
"X_RAY": [("opacity", 0.86), ("fracture", 0.78)],
|
||||
"CT": [("nodule", 0.88), ("lesion", 0.81), ("calcification", 0.74)],
|
||||
"MRI": [("lesion", 0.84), ("mass", 0.79)],
|
||||
"ULTRASOUND": [("mass", 0.80), ("lesion", 0.76)],
|
||||
}
|
||||
pairs = catalog.get(study_type.upper(), [("object", 0.75)])
|
||||
n = 1 if rng.random() < 0.35 else 2
|
||||
chosen = pairs[:n] if len(pairs) >= n else pairs
|
||||
detections: list[Detection] = []
|
||||
for i, (label, base_conf) in enumerate(chosen):
|
||||
bw = int(w * rng.uniform(0.12, 0.28))
|
||||
bh = int(h * rng.uniform(0.12, 0.28))
|
||||
x1 = int(rng.uniform(0.1, 0.65) * w)
|
||||
y1 = int(rng.uniform(0.1, 0.65) * h)
|
||||
x2 = min(w - 1, x1 + bw)
|
||||
y2 = min(h - 1, y1 + bh)
|
||||
conf = min(0.98, base_conf + rng.uniform(-0.05, 0.08) - i * 0.03)
|
||||
detections.append(
|
||||
Detection(
|
||||
label=label,
|
||||
label_zh=LABEL_ZH.get(label, label),
|
||||
confidence=round(conf, 4),
|
||||
bbox=[float(x1), float(y1), float(x2), float(y2)],
|
||||
)
|
||||
)
|
||||
return detections
|
||||
|
||||
def annotate(self, image_bgr: np.ndarray, detections: list[Detection]) -> str:
|
||||
"""返回 JPEG base64(无 data URL 前缀)。"""
|
||||
canvas = image_bgr.copy()
|
||||
for det in detections:
|
||||
x1, y1, x2, y2 = [int(v) for v in det.bbox]
|
||||
color = (40, 120, 255) if self._mode == "demo" else (46, 204, 113)
|
||||
cv2.rectangle(canvas, (x1, y1), (x2, y2), color, 2)
|
||||
text = f"{det.label_zh or det.label} {det.confidence:.0%}"
|
||||
(tw, th), _ = cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, 0.55, 1)
|
||||
cv2.rectangle(canvas, (x1, max(0, y1 - th - 8)), (x1 + tw + 6, y1), color, -1)
|
||||
cv2.putText(
|
||||
canvas,
|
||||
text,
|
||||
(x1 + 3, y1 - 5),
|
||||
cv2.FONT_HERSHEY_SIMPLEX,
|
||||
0.55,
|
||||
(255, 255, 255),
|
||||
1,
|
||||
cv2.LINE_AA,
|
||||
)
|
||||
# 角标:模式 / 权重
|
||||
badge = f"YOLO:{self._mode}"
|
||||
if self._loaded_path:
|
||||
badge += f" | {PathName(self._loaded_path)}"
|
||||
cv2.putText(
|
||||
canvas, badge, (10, 24), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (20, 20, 20), 3, cv2.LINE_AA
|
||||
)
|
||||
cv2.putText(
|
||||
canvas, badge, (10, 24), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (255, 255, 255), 1, cv2.LINE_AA
|
||||
)
|
||||
ok, buf = cv2.imencode(".jpg", canvas, [int(cv2.IMWRITE_JPEG_QUALITY), 88])
|
||||
if not ok:
|
||||
raise RuntimeError("标注图编码失败")
|
||||
return base64.b64encode(buf.tobytes()).decode("ascii")
|
||||
|
||||
|
||||
def PathName(p: str) -> str:
|
||||
from pathlib import Path
|
||||
return Path(p).name
|
||||
|
||||
|
||||
_detector: YoloDetector | None = None
|
||||
|
||||
|
||||
def get_detector() -> YoloDetector:
|
||||
global _detector
|
||||
if _detector is None:
|
||||
_detector = YoloDetector()
|
||||
return _detector
|
||||
|
||||
|
||||
def reload_detector() -> dict[str, Any]:
|
||||
det = get_detector()
|
||||
return det.reload()
|
||||
Reference in New Issue
Block a user