Files
Shixun/ai-service/app/services/yolo_detector.py
T
2026-07-31 12:27:41 +08:00

247 lines
8.4 KiB
Python

"""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()