26-7-31-1

This commit is contained in:
shuai
2026-07-31 12:27:41 +08:00
commit 9b8b8b57b3
142 changed files with 22408 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""Smart Hospital AI Service."""
+1
View File
@@ -0,0 +1 @@
"""API routers."""
+26
View File
@@ -0,0 +1,26 @@
from fastapi import APIRouter
from app.config import get_settings
from app.schemas.models import HealthResponse
from app.services.llm_client import get_llm
from app.services.monai_preprocess import monai_available
from app.services.rag_pipeline import get_rag, langchain_available
from app.services.yolo_detector import yolo_available
router = APIRouter(tags=["health"])
@router.get("/health", response_model=HealthResponse)
def health() -> HealthResponse:
settings = get_settings()
rag = get_rag()
llm = get_llm()
return HealthResponse(
status="ok",
yolo_available=yolo_available(),
monai_available=monai_available(),
langchain_available=langchain_available(),
llm_configured=llm.enabled,
demo_mode=settings.demo_mode,
knowledge_docs=rag.doc_count(),
)
+89
View File
@@ -0,0 +1,89 @@
from __future__ import annotations
import logging
from pathlib import Path
from fastapi import APIRouter, File, Form, HTTPException, UploadFile
from app.schemas.models import ImagingAnalyzeResponse
from app.services.monai_preprocess import load_image_bgr, preprocess
from app.services.report_generator import build_imaging_texts, generate_imaging_report, make_full_report
from app.schemas.models import ImagingReportRequest
from app.services.yolo_detector import get_detector
from app.services.yolo_manager import get_yolo_manager
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/imaging", tags=["imaging"])
@router.post("/analyze", response_model=ImagingAnalyzeResponse)
async def analyze_imaging(
file: UploadFile | None = File(default=None),
image_path: str | None = Form(default=None),
study_type: str = Form(default="CT"),
body_part: str = Form(default=""),
patient_summary: str = Form(default=""),
) -> ImagingAnalyzeResponse:
raw = await _read_bytes(file, image_path)
try:
image_bgr = load_image_bgr(raw)
prep = preprocess(image_bgr)
detector = get_detector()
detections = detector.detect(image_bgr, study_type=study_type)
annotated = detector.annotate(image_bgr, detections)
findings, diagnosis, recommendations, confidence = build_imaging_texts(
study_type, body_part, detections, detector.mode
)
report = generate_imaging_report(
ImagingReportRequest(
study_type=study_type,
body_part=body_part,
patient_summary=patient_summary,
preliminary_diagnosis=diagnosis,
findings=findings,
detections=detections,
confidence=confidence,
)
)
backend = prep.get("backend", "opencv")
model_version = f"yolo-{detector.mode}+{backend}+{report.model_version}"
try:
get_yolo_manager().record_inference(
mode=detector.mode,
detections=detections,
study_type=study_type,
model_version=model_version,
)
except Exception as e:
logger.warning("记录 YOLO 统计失败: %s", e)
return ImagingAnalyzeResponse(
detections=detections,
annotated_image_base64=annotated,
preliminary_diagnosis=report.impression or diagnosis,
findings=report.findings or findings,
recommendations=report.recommendations or recommendations,
confidence=confidence,
model_version=model_version,
mode=detector.mode, # type: ignore[arg-type]
full_report=report.full_report
or make_full_report(study_type, body_part, findings, diagnosis, recommendations, patient_summary),
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from e
except Exception as e:
logger.exception("影像分析失败")
raise HTTPException(status_code=500, detail=f"影像分析失败: {e}") from e
async def _read_bytes(file: UploadFile | None, image_path: str | None) -> bytes:
if file is not None:
data = await file.read()
if not data:
raise HTTPException(status_code=400, detail="上传文件为空")
return data
if image_path:
path = Path(image_path)
if not path.is_file():
raise HTTPException(status_code=400, detail=f"影像路径不存在: {image_path}")
return path.read_bytes()
raise HTTPException(status_code=400, detail="请提供 file 或 image_path")
+70
View File
@@ -0,0 +1,70 @@
"""管理端下发的 LLM 运行时配置。"""
from __future__ import annotations
import logging
from fastapi import APIRouter
from pydantic import BaseModel, Field
from app.config import get_settings
from app.services.llm_client import get_llm
from app.services.llm_runtime import get_runtime_llm, update_runtime_llm
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/llm", tags=["llm"])
class LlmConfigUpdate(BaseModel):
enabled: bool | None = True
api_base_url: str | None = Field(default=None, alias="api_base_url")
api_key: str | None = None
model: str | None = None
temperature: float | None = None
# 兼容 camelCase(Spring 默认可能发 apiBaseUrl)
apiBaseUrl: str | None = None
apiKey: str | None = None
class Config:
populate_by_name = True
@router.get("/config")
def get_llm_config() -> dict:
return get_llm().info()
@router.post("/config")
def set_llm_config(body: LlmConfigUpdate) -> dict:
base = body.api_base_url or body.apiBaseUrl
key = body.api_key if body.api_key is not None else body.apiKey
# 注意:key 为 None 表示本次不改密钥;空字符串表示清空运行时密钥
rt = update_runtime_llm(
enabled=body.enabled,
api_base_url=base,
api_key=key,
model=body.model,
temperature=body.temperature,
)
info = get_llm().info()
logger.info(
"已更新运行时 LLM 配置 enabled=%s model=%s base=%s key=%s source=%s",
info.get("enabled"),
info.get("model"),
info.get("api_base_url"),
"yes" if info.get("api_key_configured") else "no",
rt.source,
)
return {"ok": True, **info}
@router.get("/status")
def llm_status() -> dict:
settings = get_settings()
info = get_llm().info()
rt = get_runtime_llm()
return {
**info,
"env_key_configured": bool(settings.llm_api_key and settings.llm_api_key.strip()),
"runtime_enabled_flag": rt.enabled,
}
+49
View File
@@ -0,0 +1,49 @@
from fastapi import APIRouter
from pydantic import BaseModel, Field
from app.schemas.models import RagQueryRequest, RagQueryResponse, SourceRef
from app.services.rag_pipeline import get_rag
router = APIRouter(prefix="/rag", tags=["rag"])
class IngestRequest(BaseModel):
title: str
content: str
category: str = "自定义"
class IngestResponse(BaseModel):
ok: bool = True
chunks: int = 0
@router.post("/query", response_model=RagQueryResponse)
def rag_query(body: RagQueryRequest) -> RagQueryResponse:
rag = get_rag()
result = rag.query(body.query, top_k=body.top_k, extra_context=body.context or "")
raw_sources = result.get("sources") or []
sources: list[SourceRef] = []
for s in raw_sources:
if isinstance(s, SourceRef):
sources.append(s)
elif isinstance(s, dict):
sources.append(SourceRef(**s))
return RagQueryResponse(
answer=result.get("answer") or "",
sources=sources,
engine=result.get("engine") or "langchain-rag",
)
@router.post("/ingest", response_model=IngestResponse)
def rag_ingest(body: IngestRequest) -> IngestResponse:
rag = get_rag()
rag.ingest_text(body.title, body.content, body.category)
return IngestResponse(ok=True, chunks=rag.doc_count())
@router.get("/stats")
def rag_stats() -> dict:
rag = get_rag()
return {"chunks": rag.doc_count()}
+21
View File
@@ -0,0 +1,21 @@
from fastapi import APIRouter
from app.schemas.models import (
DecisionRequest,
DecisionResponse,
ImagingReportRequest,
ImagingReportResponse,
)
from app.services.report_generator import generate_decision, generate_imaging_report
router = APIRouter(prefix="/report", tags=["report"])
@router.post("/imaging", response_model=ImagingReportResponse)
def report_imaging(body: ImagingReportRequest) -> ImagingReportResponse:
return generate_imaging_report(body)
@router.post("/decision", response_model=DecisionResponse)
def report_decision(body: DecisionRequest) -> DecisionResponse:
return generate_decision(body)
+135
View File
@@ -0,0 +1,135 @@
from __future__ import annotations
import logging
from fastapi import APIRouter, File, HTTPException, UploadFile
from pydantic import BaseModel, Field
from app.services.yolo_detector import get_detector, reload_detector, yolo_available
from app.services.yolo_manager import get_yolo_manager
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/yolo", tags=["yolo"])
class ActivateRequest(BaseModel):
name: str = Field(..., description="权重文件名")
demo_mode: str | None = Field(default=None, description="demo|real|auto")
class DemoModeRequest(BaseModel):
demo_mode: str = Field(..., description="demo|real|auto")
@router.get("/status")
def yolo_status() -> dict:
mgr = get_yolo_manager()
det = get_detector()
return mgr.status(det.info())
@router.get("/weights")
def list_weights() -> dict:
mgr = get_yolo_manager()
return {"items": mgr.list_weights(), "count": len(mgr.list_weights())}
@router.post("/weights/upload")
async def upload_weight(file: UploadFile = File(...)) -> dict:
if not file.filename:
raise HTTPException(400, "缺少文件名")
raw = await file.read()
try:
item = get_yolo_manager().save_upload(file.filename, raw)
return {"ok": True, "weight": item}
except ValueError as e:
raise HTTPException(400, str(e)) from e
except Exception as e:
logger.exception("上传权重失败")
raise HTTPException(500, f"上传失败: {e}") from e
@router.post("/weights/activate")
def activate_weight(body: ActivateRequest) -> dict:
try:
status = get_yolo_manager().activate(body.name, body.demo_mode)
info = reload_detector()
status.update(info)
status["ok"] = True
return status
except FileNotFoundError as e:
raise HTTPException(404, str(e)) from e
except ValueError as e:
raise HTTPException(400, str(e)) from e
@router.post("/weights/deactivate")
def deactivate_weight() -> dict:
status = get_yolo_manager().deactivate()
info = reload_detector()
status.update(info)
status["ok"] = True
return status
@router.delete("/weights/{name}")
def delete_weight(name: str) -> dict:
"""逻辑删除:隐藏权重,本地 .pt 文件仍保留在 data/weights。"""
try:
get_yolo_manager().delete_weight(name)
reload_detector()
return {
"ok": True,
"name": name,
"soft_delete": True,
"message": "已逻辑删除(本地文件保留,仅从列表隐藏)",
}
except FileNotFoundError as e:
raise HTTPException(404, str(e)) from e
except ValueError as e:
raise HTTPException(400, str(e)) from e
@router.post("/weights/restore")
def restore_weight(body: ActivateRequest) -> dict:
"""恢复逻辑删除的权重。"""
try:
status = get_yolo_manager().restore_weight(body.name)
info = reload_detector()
status.update(info)
status["ok"] = True
return status
except FileNotFoundError as e:
raise HTTPException(404, str(e)) from e
@router.post("/mode")
def set_mode(body: DemoModeRequest) -> dict:
try:
status = get_yolo_manager().set_demo_mode(body.demo_mode)
info = reload_detector()
status.update(info)
status["ok"] = True
return status
except ValueError as e:
raise HTTPException(400, str(e)) from e
@router.get("/visualization")
def visualization() -> dict:
return get_yolo_manager().visualization()
@router.post("/stats/reset")
def reset_stats() -> dict:
get_yolo_manager().reset_stats()
return {"ok": True}
@router.get("/capability")
def capability() -> dict:
return {
"ultralytics": yolo_available(),
"detector": get_detector().info(),
"manager": get_yolo_manager().status(),
}
+63
View File
@@ -0,0 +1,63 @@
from __future__ import annotations
from functools import lru_cache
from pathlib import Path
from typing import Literal
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
extra="ignore",
)
ai_host: str = "0.0.0.0"
ai_port: int = 8001
demo_mode: Literal["demo", "real", "auto"] = "auto"
yolo_weights: str = ""
llm_base_url: str = "https://api.deepseek.com"
llm_api_key: str = ""
llm_model: str = "deepseek-chat"
llm_temperature: float = 0.3
knowledge_dir: str = "./app/knowledge"
vector_dir: str = "./data/vectorstore"
@property
def root_dir(self) -> Path:
return Path(__file__).resolve().parent.parent
def resolve_path(self, value: str) -> Path:
p = Path(value)
if p.is_absolute():
return p
return (self.root_dir / p).resolve()
@property
def knowledge_path(self) -> Path:
return self.resolve_path(self.knowledge_dir)
@property
def vector_path(self) -> Path:
return self.resolve_path(self.vector_dir)
@property
def yolo_weights_path(self) -> Path | None:
if not self.yolo_weights or not self.yolo_weights.strip():
return None
path = self.resolve_path(self.yolo_weights.strip())
return path if path.is_file() else None
@property
def llm_enabled(self) -> bool:
return bool(self.llm_api_key and self.llm_api_key.strip())
@lru_cache
def get_settings() -> Settings:
return Settings()
@@ -0,0 +1,19 @@
# 胆囊结石辅助管理要点
category: 消化外科
## 表现
右上腹痛、油腻饮食后加重,可伴恶心;超声是首选检查。
## 处理
- 无症状结石可观察
- 症状性结石评估腹腔镜胆囊切除指征
- 合并胆管炎/胰腺炎需紧急处理
## 护理与饮食
- 低脂饮食
- 观察腹痛、黄疸、发热
- 术后早期活动
## 注意
实训演示用,非临床处方依据。
+24
View File
@@ -0,0 +1,24 @@
# 2 型糖尿病辅助管理要点
category: 内分泌
## 诊断要点
空腹血糖、OGTT 或 HbA1c 达到诊断阈值;需排除 1 型及其他特殊类型。
## 综合管理
- 医学营养治疗与运动
- 血糖自我监测
- 个体化降糖药物(如二甲双胍等,注意禁忌)
- 血压、血脂与抗血小板综合干预
## 并发症筛查
- 视网膜病变、肾病、神经病变、足病
- 心血管风险评估
## 护理与随访
- 低血糖识别与处理
- 足部护理
- 定期复查血糖与 HbA1c
## 注意
本资料用于教学演示,临床决策需结合指南与患者情况。
+20
View File
@@ -0,0 +1,20 @@
# 骨折影像与处置要点(示意)
category: 骨科
## 影像
- X 线是基础检查;复杂部位可 CT/MRI。
- 描述骨折部位、类型、移位、关节受累。
## 处理原则
- 复位、固定、功能锻炼
- 开放伤注意清创抗感染
- 评估血管神经损伤
## 护理与随访
- 患肢肿胀、血运、感觉观察
- 疼痛管理与防深静脉血栓
- 按医嘱复查 X 光评估愈合
## 注意
仅供实训演示,不作为临床操作依据。
+17
View File
@@ -0,0 +1,17 @@
# 电子病历与辅助决策书写规范(示意)
category: 病历质控
## 病历要素
主诉、现病史、体格检查、辅助检查、诊断、治疗计划、用药、随访。
## AI 辅助使用原则
- AI 输出仅供参考,必须经医师审核修改后入档。
- 引用知识库内容时注明来源类别。
- 避免将模型幻觉内容写入正式病历。
## 护理建议常见结构
病情观察、用药护理、生活指导、心理支持、健康宣教。
## 随访计划
复诊时间、复查项目、危险症状返院指征。
+24
View File
@@ -0,0 +1,24 @@
# 原发性高血压辅助管理要点
category: 心血管
## 诊断
非同日三次诊室血压 ≥140/90 mmHg 可诊断高血压;鼓励家庭血压与动态血压监测。
## 非药物治疗
- 限盐(<5 g/日)、DASH 饮食
- 规律有氧运动、控制体重
- 戒烟限酒、减少精神压力
## 药物治疗(示意)
- 常用类别:ACEI/ARB、CCB、利尿剂、β 受体阻滞剂等。
- 合并糖尿病、慢性肾病等需个体化选药。
- 注意体位性低血压与电解质紊乱。
## 护理与随访
- 规范测量并记录血压
- 提高服药依从性宣教
- 评估靶器官损害,定期复诊
## 注意
用药方案必须由执业医师开具,本文仅供实训演示。
+22
View File
@@ -0,0 +1,22 @@
# 颈椎/腰椎间盘突出 MRI 辅助解读要点
category: 骨科/影像
## 常见表现
- 椎间盘信号改变、突出压迫硬膜囊或神经根
- 可伴椎管狭窄、黄韧带肥厚
## 临床相关
- 颈肩痛、上肢放射痛、腰痛、下肢放射痛、麻木
- 需与脊髓病、肿瘤、感染鉴别
## 治疗路径(示意)
- 多数可先保守:休息、理疗、药物、功能锻炼
- 神经功能障碍加重或马尾症状需紧急评估手术指征
## 随访
- 症状变化记录
- 必要时复查 MRI
## 注意
教学演示资料,诊断以影像科与临床医师为准。
+23
View File
@@ -0,0 +1,23 @@
# 社区获得性肺炎辅助诊疗要点
category: 呼吸内科
## 概述
社区获得性肺炎(CAP)是指在医院外罹患的肺实质感染,常见症状包括发热、咳嗽、咳痰、胸痛与呼吸困难。
## 影像表现
- X 线/CT:片状、斑片状浸润影或实变,可伴空气支气管征。
- 需与肺结核、肿瘤、肺水肿等鉴别。
## 治疗原则
- 评估 CURB-65 或 PSI 严重程度分层。
- 经验性抗感染覆盖常见病原(肺炎链球菌等),并根据培养结果调整。
- 支持治疗:氧疗、补液、退热、止咳化痰。
## 护理与随访
- 监测体温、呼吸频率、血氧饱和度。
- 鼓励有效咳嗽与体位引流。
- 治疗后 3–5 天评估疗效;必要时复查胸片。
## 注意
本资料仅供教学演示与辅助决策参考,不能替代临床指南与医师判断。
@@ -0,0 +1,23 @@
# 肺结节随访与管理要点
category: 呼吸/胸外
## 定义
肺结节通常指直径 ≤3 cm 的圆形或类圆形 dens 灶,需结合大小、形态、密度及随访变化综合判断。
## 影像关注点
- 大小与倍增时间
- 边缘(毛刺、分叶)、钙化、空泡、胸膜牵拉
- 实性 / 部分实性 / 磨玻璃
## 管理建议(示意)
- 微小结节:定期 CT 随访。
- 可疑恶性特征:建议专科评估,必要时 PET、活检或多学科讨论。
- 戒烟是重要干预措施。
## 护理与宣教
- 缓解患者焦虑,说明随访意义。
- 出现咯血、胸痛、进行性气促及时就诊。
## 注意
本资料仅供教学演示,实际随访间隔需按最新指南与个体情况确定。
+75
View File
@@ -0,0 +1,75 @@
from __future__ import annotations
import logging
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.api import health, imaging, llm_config, rag, report, yolo
from app.config import get_settings
from app.services.llm_client import get_llm
from app.services.rag_pipeline import get_rag
from app.services.yolo_detector import get_detector
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
)
logger = logging.getLogger("ai-service")
settings = get_settings()
app = FastAPI(
title="Smart Hospital AI Service",
description="YOLO/MONAI 影像识别 + LangChain RAG + 报告/决策生成(实训演示)",
version="1.0.0",
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(health.router)
app.include_router(imaging.router)
app.include_router(rag.router)
app.include_router(report.router)
app.include_router(yolo.router)
app.include_router(llm_config.router)
@app.on_event("startup")
def on_startup() -> None:
rag_pipe = get_rag()
det_info = get_detector().info()
llm_info = get_llm().info()
logger.info(
"AI 服务启动 port=%s knowledge_chunks=%s llm=%s yolo_mode=%s",
settings.ai_port,
rag_pipe.doc_count(),
llm_info.get("enabled"),
det_info.get("runtime_mode"),
)
@app.get("/")
def root() -> dict:
return {
"service": "smart-hospital-ai",
"docs": "/docs",
"health": "/health",
}
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"app.main:app",
host=settings.ai_host,
port=settings.ai_port,
reload=False,
)
+27
View File
@@ -0,0 +1,27 @@
from .models import (
Detection,
ImagingAnalyzeResponse,
ImagingReportRequest,
ImagingReportResponse,
DecisionRequest,
DecisionResponse,
RiskItem,
SourceRef,
RagQueryRequest,
RagQueryResponse,
HealthResponse,
)
__all__ = [
"Detection",
"ImagingAnalyzeResponse",
"ImagingReportRequest",
"ImagingReportResponse",
"DecisionRequest",
"DecisionResponse",
"RiskItem",
"SourceRef",
"RagQueryRequest",
"RagQueryResponse",
"HealthResponse",
]
+108
View File
@@ -0,0 +1,108 @@
from __future__ import annotations
from typing import Any, Literal, Optional
from pydantic import BaseModel, Field
class Detection(BaseModel):
label: str
confidence: float
bbox: list[float] = Field(description="[x1, y1, x2, y2] 像素坐标")
label_zh: Optional[str] = None
class ImagingAnalyzeResponse(BaseModel):
detections: list[Detection] = []
annotated_image_base64: Optional[str] = None
preliminary_diagnosis: str
findings: str
recommendations: str
confidence: float
model_version: str
mode: Literal["demo", "real"] = "demo"
full_report: Optional[str] = None
disclaimer: str = "本结果仅供辅助决策,不能替代执业医师诊断。"
class ImagingReportRequest(BaseModel):
study_type: str = "CT"
body_part: str = ""
patient_summary: str = ""
preliminary_diagnosis: str = ""
findings: str = ""
detections: list[Detection] = []
confidence: float = 0.85
class ImagingReportResponse(BaseModel):
findings: str
impression: str
recommendations: str
full_report: str
model_version: str = "report-v1"
class PatientInfo(BaseModel):
age: Optional[int] = None
gender: Optional[str] = None
name: Optional[str] = None
class DecisionRequest(BaseModel):
chief_complaint: str = ""
history: str = ""
exam_findings: str = ""
diagnosis: str = ""
medications: str = ""
imaging_summary: str = ""
patient: Optional[PatientInfo] = None
class RiskItem(BaseModel):
type: str
description: str
level: str = "中"
confidence: float = 0.8
class SourceRef(BaseModel):
title: str
snippet: str = ""
category: str = ""
score: float = 0.0
class DecisionResponse(BaseModel):
treatment_suggestions: list[dict[str, Any]] = []
medication_suggestions: list[dict[str, Any]] = []
nursing_advice: list[str] = []
follow_up_plan: list[str] = []
risks: list[RiskItem] = []
conflicts: list[str] = []
sources: list[SourceRef] = []
full_text: str = ""
engine: str = "fastapi-rag"
disclaimer: str = "本结果仅供辅助决策,不能替代执业医师诊断。"
class RagQueryRequest(BaseModel):
query: str
top_k: int = 4
context: str = ""
class RagQueryResponse(BaseModel):
answer: str
sources: list[SourceRef] = []
engine: str = "langchain-rag"
class HealthResponse(BaseModel):
status: str = "ok"
yolo_available: bool = False
monai_available: bool = False
langchain_available: bool = False
llm_configured: bool = False
demo_mode: str = "auto"
knowledge_docs: int = 0
+1
View File
@@ -0,0 +1 @@
"""AI service implementations."""
+142
View File
@@ -0,0 +1,142 @@
"""OpenAI 兼容 LLM 客户端(DeepSeek / Qwen)。支持 .env + 管理端运行时覆盖。"""
from __future__ import annotations
import json
import logging
import re
from typing import Any
import httpx
from app.config import Settings, get_settings
from app.services.llm_runtime import get_runtime_llm
logger = logging.getLogger(__name__)
class LlmClient:
def __init__(self, settings: Settings | None = None):
self.settings = settings or get_settings()
def _effective_key(self) -> str:
rt = get_runtime_llm()
if rt.api_key is not None:
return rt.api_key.strip()
return (self.settings.llm_api_key or "").strip()
def _effective_base_url(self) -> str:
rt = get_runtime_llm()
if rt.api_base_url:
return rt.api_base_url.rstrip("/")
return (self.settings.llm_base_url or "https://api.deepseek.com").rstrip("/")
def _effective_model(self) -> str:
rt = get_runtime_llm()
if rt.model:
return rt.model
return self.settings.llm_model or "deepseek-chat"
def _effective_temperature(self, override: float | None = None) -> float:
if override is not None:
return override
rt = get_runtime_llm()
if rt.temperature is not None:
return float(rt.temperature)
return float(self.settings.llm_temperature)
@property
def enabled(self) -> bool:
"""管理端 enabled=false 强制关闭;否则有可用 API Key 即启用。"""
key = self._effective_key()
if not key:
return False
rt = get_runtime_llm()
if rt.enabled is False:
return False
if rt.enabled is True:
return True
# 未下发 enabled 时:有 key(env 或 runtime)即视为可用
return True
def info(self) -> dict[str, Any]:
rt = get_runtime_llm()
return {
"enabled": self.enabled,
"api_key_configured": bool(self._effective_key()),
"api_base_url": self._effective_base_url(),
"model": self._effective_model(),
"temperature": self._effective_temperature(),
"source": rt.source if (rt.api_key or rt.enabled is not None) else "env",
}
def chat(self, messages: list[dict[str, str]], temperature: float | None = None) -> str:
if not self.enabled:
raise RuntimeError("未配置 LLM(请在管理端「AI 配置」启用并填写 API Key,或设置 ai-service/.env 的 LLM_API_KEY)")
base = self._effective_base_url()
if base.endswith("/v1"):
url = base + "/chat/completions"
else:
url = base + "/v1/chat/completions"
payload = {
"model": self._effective_model(),
"temperature": self._effective_temperature(temperature),
"messages": messages,
}
headers = {
"Authorization": f"Bearer {self._effective_key()}",
"Content-Type": "application/json",
}
with httpx.Client(timeout=90.0) as client:
resp = client.post(url, headers=headers, json=payload)
if resp.status_code >= 400:
raise RuntimeError(f"LLM HTTP {resp.status_code}: {resp.text[:300]}")
data = resp.json()
content = (
data.get("choices", [{}])[0]
.get("message", {})
.get("content", "")
)
if not content:
raise RuntimeError("LLM 返回空内容")
return content.strip()
def chat_json(self, messages: list[dict[str, str]]) -> dict[str, Any]:
text = self.chat(messages, temperature=0.2)
return extract_json(text)
def extract_json(text: str) -> dict[str, Any]:
text = text.strip()
try:
return json.loads(text)
except json.JSONDecodeError:
pass
fence = re.search(r"```(?:json)?\s*([\s\S]*?)```", text)
if fence:
try:
return json.loads(fence.group(1).strip())
except json.JSONDecodeError:
pass
start, end = text.find("{"), text.rfind("}")
if start >= 0 and end > start:
try:
return json.loads(text[start : end + 1])
except json.JSONDecodeError:
pass
raise ValueError("无法从模型输出解析 JSON")
_llm: LlmClient | None = None
def get_llm() -> LlmClient:
global _llm
if _llm is None:
_llm = LlmClient()
return _llm
def reset_llm_client() -> None:
"""测试或热更新后可重置单例(配置本身已从 runtime 动态读取,一般无需调用)。"""
global _llm
_llm = None
+84
View File
@@ -0,0 +1,84 @@
"""运行时 LLM 配置:可由业务后端(管理端 AI 配置)动态下发,覆盖 .env。"""
from __future__ import annotations
from dataclasses import dataclass, field
from threading import RLock
from typing import Any
@dataclass
class RuntimeLlmConfig:
"""enabled=None 表示未由管理端覆盖,沿用 .env。"""
enabled: bool | None = None
api_base_url: str | None = None
api_key: str | None = None
model: str | None = None
temperature: float | None = None
source: str = "env"
_lock = RLock()
_runtime = RuntimeLlmConfig()
def get_runtime_llm() -> RuntimeLlmConfig:
with _lock:
return RuntimeLlmConfig(
enabled=_runtime.enabled,
api_base_url=_runtime.api_base_url,
api_key=_runtime.api_key,
model=_runtime.model,
temperature=_runtime.temperature,
source=_runtime.source,
)
def update_runtime_llm(
*,
enabled: bool | None = None,
api_base_url: str | None = None,
api_key: str | None = None,
model: str | None = None,
temperature: float | None = None,
) -> RuntimeLlmConfig:
"""
更新运行时配置。
- api_key 为 None:不改密钥
- api_key 为非空字符串:覆盖
- api_key 为 "":清空运行时密钥(回退 .env)
"""
with _lock:
if enabled is not None:
_runtime.enabled = bool(enabled)
if api_base_url is not None and api_base_url.strip():
_runtime.api_base_url = api_base_url.strip()
if api_key is not None:
_runtime.api_key = api_key.strip() if api_key.strip() else None
if model is not None and model.strip():
_runtime.model = model.strip()
if temperature is not None:
t = float(temperature)
_runtime.temperature = max(0.0, min(2.0, t))
_runtime.source = "runtime"
return get_runtime_llm()
def runtime_status(env_key_configured: bool, env_base: str, env_model: str) -> dict[str, Any]:
rt = get_runtime_llm()
key_ok = bool(rt.api_key) or env_key_configured
enabled = False
if key_ok:
if rt.enabled is None:
enabled = env_key_configured or bool(rt.api_key)
else:
enabled = bool(rt.enabled)
return {
"enabled": enabled,
"configured": key_ok,
"api_key_configured": key_ok,
"api_base_url": rt.api_base_url or env_base,
"model": rt.model or env_model,
"temperature": rt.temperature,
"source": rt.source if (rt.api_key or rt.enabled is not None or rt.api_base_url) else "env",
}
@@ -0,0 +1,78 @@
"""医学影像预处理:优先 MONAI,失败则用 OpenCV/Pillow 降级。"""
from __future__ import annotations
import logging
from typing import Any
import cv2
import numpy as np
logger = logging.getLogger(__name__)
_MONAI_OK = False
try:
import monai # noqa: F401
from monai.transforms import Compose, ScaleIntensity, Resize
_MONAI_OK = True
except Exception: # pragma: no cover
_MONAI_OK = False
logger.info("MONAI 未安装,使用 OpenCV 预处理管线")
def monai_available() -> bool:
return _MONAI_OK
def load_image_bgr(image_bytes: bytes) -> np.ndarray:
arr = np.frombuffer(image_bytes, dtype=np.uint8)
img = cv2.imdecode(arr, cv2.IMREAD_COLOR)
if img is None:
raise ValueError("无法解码影像文件,请上传常见图片格式(jpg/png 等)")
return img
def preprocess(image_bgr: np.ndarray, target_size: int = 640) -> dict[str, Any]:
"""
返回:
- image_bgr: 原始 BGR
- image_rgb: RGB
- tensor_like: 归一化后的 float32 CHW(MONAI 或 numpy 模拟)
- meta: 尺寸信息
"""
h, w = image_bgr.shape[:2]
rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
if _MONAI_OK:
try:
# 灰度/三通道统一为 CHW float,再经 MONAI ScaleIntensity + Resize
chw = np.transpose(rgb.astype(np.float32) / 255.0, (2, 0, 1))
transforms = Compose(
[
ScaleIntensity(minv=0.0, maxv=1.0),
Resize(spatial_size=(target_size, target_size), mode="bilinear"),
]
)
tensor = transforms(chw)
if hasattr(tensor, "numpy"):
tensor = tensor.numpy()
return {
"image_bgr": image_bgr,
"image_rgb": rgb,
"tensor_like": np.asarray(tensor),
"backend": "monai",
"meta": {"orig_h": h, "orig_w": w, "target": target_size},
}
except Exception as e: # pragma: no cover
logger.warning("MONAI 预处理失败,降级 OpenCV: %s", e)
# OpenCV 降级:resize + normalize
resized = cv2.resize(rgb, (target_size, target_size), interpolation=cv2.INTER_LINEAR)
tensor = np.transpose(resized.astype(np.float32) / 255.0, (2, 0, 1))
return {
"image_bgr": image_bgr,
"image_rgb": rgb,
"tensor_like": tensor,
"backend": "opencv",
"meta": {"orig_h": h, "orig_w": w, "target": target_size},
}
+242
View File
@@ -0,0 +1,242 @@
"""LangChain 风格 RAG:文档切分 + 关键词检索 + 可选 LLM 生成。"""
from __future__ import annotations
import logging
import re
from dataclasses import dataclass
from typing import Any
from app.config import Settings, get_settings
from app.schemas.models import SourceRef
from app.services.llm_client import get_llm
logger = logging.getLogger(__name__)
_LC_OK = False
try:
from langchain_core.documents import Document
from langchain_text_splitters import RecursiveCharacterTextSplitter
_LC_OK = True
except Exception: # pragma: no cover
Document = None # type: ignore
RecursiveCharacterTextSplitter = None # type: ignore
logger.info("LangChain 未完全安装,将使用内置简易检索")
def langchain_available() -> bool:
return _LC_OK
@dataclass
class Chunk:
title: str
category: str
content: str
source: str
class RagPipeline:
def __init__(self, settings: Settings | None = None):
self.settings = settings or get_settings()
self.chunks: list[Chunk] = []
self._load_knowledge()
def _load_knowledge(self) -> None:
path = self.settings.knowledge_path
path.mkdir(parents=True, exist_ok=True)
files = sorted(list(path.glob("*.md")) + list(path.glob("*.txt")))
raw_docs: list[tuple[str, str, str]] = []
for f in files:
try:
text = f.read_text(encoding="utf-8")
except Exception:
continue
title, category, body = self._parse_doc(f.stem, text)
raw_docs.append((title, category, body))
if not raw_docs:
logger.warning("知识库目录为空: %s", path)
self.chunks = []
return
if _LC_OK:
splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=80)
for title, category, body in raw_docs:
docs = splitter.split_documents(
[Document(page_content=body, metadata={"title": title, "category": category})]
)
for d in docs:
self.chunks.append(
Chunk(
title=title,
category=category,
content=d.page_content,
source=title,
)
)
else:
for title, category, body in raw_docs:
for part in self._simple_split(body, 500):
self.chunks.append(
Chunk(title=title, category=category, content=part, source=title)
)
logger.info("知识库已加载 %d 个文档片段", len(self.chunks))
@staticmethod
def _parse_doc(stem: str, text: str) -> tuple[str, str, str]:
title = stem
category = "临床指南"
body = text.strip()
lines = body.splitlines()
if lines and lines[0].startswith("#"):
title = lines[0].lstrip("#").strip() or stem
body = "\n".join(lines[1:]).strip()
m = re.search(r"category:\s*(.+)", body, re.I)
if m:
category = m.group(1).strip()
return title, category, body
@staticmethod
def _simple_split(text: str, size: int) -> list[str]:
if len(text) <= size:
return [text]
parts: list[str] = []
i = 0
while i < len(text):
parts.append(text[i : i + size])
i += max(1, size - 50)
return parts
def reload(self) -> int:
self.chunks = []
self._load_knowledge()
return len(self.chunks)
def doc_count(self) -> int:
return len(self.chunks)
def retrieve(self, query: str, top_k: int = 4) -> list[SourceRef]:
if not query or not self.chunks:
return []
tokens = self._tokenize(query)
scored: list[tuple[float, Chunk]] = []
q = query.lower()
for ch in self.chunks:
score = self._score(ch, tokens, q)
if score > 0:
scored.append((score, ch))
scored.sort(key=lambda x: x[0], reverse=True)
results: list[SourceRef] = []
seen: set[tuple[str, str]] = set()
for score, ch in scored:
key = (ch.title, ch.content[:80])
if key in seen:
continue
seen.add(key)
results.append(
SourceRef(
title=ch.title,
category=ch.category,
snippet=ch.content[:220].replace("\n", " "),
score=round(score, 2),
)
)
if len(results) >= top_k:
break
return results
def build_context(self, sources: list[SourceRef]) -> str:
if not sources:
return ""
parts = []
for i, s in enumerate(sources, 1):
parts.append(f"【资料{i}】{s.title}\n{s.snippet}")
return "\n\n".join(parts)
def query(self, query: str, top_k: int = 4, extra_context: str = "") -> dict[str, Any]:
sources = self.retrieve(query, top_k=top_k)
context = self.build_context(sources)
if extra_context:
context = (extra_context.strip() + "\n\n" + context).strip()
llm = get_llm()
if llm.enabled and (context or query):
try:
answer = llm.chat(
[
{
"role": "system",
"content": (
"你是医院临床辅助决策助手。请仅依据给定资料与问题作答,"
"语言专业简洁,并提醒需医师审核。不要编造未提供的检查数据。"
),
},
{
"role": "user",
"content": f"参考资料:\n{context or '(无)'}\n\n问题:{query}",
},
]
)
return {"answer": answer, "sources": sources, "engine": "langchain-rag+llm"}
except Exception as e:
logger.warning("RAG LLM 失败: %s", e)
if sources:
answer = (
f"基于本地知识库检索(LangChain 文档切分 + 关键词排序),与「{query}」相关的要点:\n"
+ "\n".join(f"- {s.title}:{s.snippet[:120]}" for s in sources)
+ "\n\n(未配置大模型 API Key 或调用失败时展示检索摘要,仅供参考。)"
)
return {"answer": answer, "sources": sources, "engine": "langchain-rag-local"}
return {
"answer": f"知识库中未检索到与「{query}」高度相关的条目,建议补充临床指南或完善病历描述。",
"sources": [],
"engine": "langchain-rag-empty",
}
def ingest_text(self, title: str, content: str, category: str = "自定义") -> None:
path = self.settings.knowledge_path
path.mkdir(parents=True, exist_ok=True)
safe = re.sub(r"[^\w\u4e00-\u9fff\-]+", "_", title)[:40] or "doc"
file_path = path / f"{safe}.md"
file_path.write_text(
f"# {title}\n\ncategory: {category}\n\n{content}\n",
encoding="utf-8",
)
self.reload()
@staticmethod
def _tokenize(text: str) -> list[str]:
parts = re.split(r"[\s,,。.!!??;;::、/\\|_\-—()()\[\]{}]+", text.lower())
return [p for p in parts if len(p) >= 2]
@staticmethod
def _score(ch: Chunk, tokens: list[str], q: str) -> float:
title = ch.title.lower()
content = ch.content.lower()
score = 0.0
if q and q in title:
score += 30
if q and q in content:
score += 15
for t in tokens:
if t in title:
score += 8
if t in content:
score += 3
if t in ch.category.lower():
score += 2
return score
_rag: RagPipeline | None = None
def get_rag() -> RagPipeline:
global _rag
if _rag is None:
_rag = RagPipeline()
return _rag
+417
View File
@@ -0,0 +1,417 @@
"""影像报告与临床决策建议生成。"""
from __future__ import annotations
import logging
import re
from typing import Any
from app.schemas.models import (
DecisionRequest,
DecisionResponse,
Detection,
ImagingReportRequest,
ImagingReportResponse,
RiskItem,
SourceRef,
)
from app.services.llm_client import get_llm
from app.services.rag_pipeline import get_rag
logger = logging.getLogger(__name__)
STUDY_LABEL = {
"X_RAY": "X 光",
"CT": "CT",
"MRI": "MRI",
"ULTRASOUND": "超声",
}
def build_imaging_texts(
study_type: str,
body_part: str,
detections: list[Detection],
mode: str,
) -> tuple[str, str, str, float]:
"""返回 findings, diagnosis, recommendations, confidence。"""
st = STUDY_LABEL.get(study_type.upper(), study_type)
part = body_part or "相关部位"
if not detections:
findings = f"{st}检查({part}):影像质量可评估,未见明确异常密度/信号灶。"
diagnosis = f"{part}{st}未见明显异常"
rec = "建议结合临床,必要时复查或进一步检查。"
return findings, diagnosis, rec, 0.82
lines = []
for d in detections:
name = d.label_zh or d.label
lines.append(
f"- 可见{name}样改变,框选区域约 ({int(d.bbox[0])},{int(d.bbox[1])})-"
f"({int(d.bbox[2])},{int(d.bbox[3])}),模型置信度 {d.confidence:.0%}"
)
findings = f"{st}检查({part})AI 辅助读片所见:\n" + "\n".join(lines)
top = max(detections, key=lambda x: x.confidence)
diagnosis = f"{part}可疑{top.label_zh or top.label},建议专科医师复核"
rec = _rec_for_label(top.label)
conf = sum(d.confidence for d in detections) / len(detections)
if mode == "demo":
findings += "\n(演示模式:检测框由 YOLO 演示引擎生成,非临床验证模型输出)"
return findings, diagnosis, rec, round(min(0.98, conf), 4)
def _rec_for_label(label: str) -> str:
mapping = {
"opacity": "建议结合血常规/炎症指标,必要时抗感染治疗并短期复查胸片。",
"nodule": "建议按结节指南分层管理,3 个月后复查 CT,必要时多学科会诊。",
"fracture": "建议骨科评估,必要时制动/固定,复查局部 X 光。",
"effusion": "建议评估积液性质,必要时穿刺或超声随访。",
"lesion": "建议结合临床与实验室检查,必要时增强扫描或专科转诊。",
"mass": "建议进一步定性检查,排除占位性病变,及时专科就诊。",
"calcification": "多为良性钙化可能,建议定期随访观察。",
}
return mapping.get(label, "建议专科医师综合临床资料判读,制定个体化方案。")
def _normalize_multiline(text: str) -> str:
"""把挤成一段的长文尽量拆成可读多行(句号/分号后换行,编号建议分行)。"""
if not text:
return ""
s = str(text).strip()
# 已有明显换行则只做空白整理
if "\n" in s and s.count("\n") >= 2:
return "\n".join(line.strip() for line in s.splitlines() if line.strip())
# 编号建议:1. / 1、 / (1) 前换行
s = re.sub(r"(?<![.\d])\s*([((]?\d+[)).、])\s*", r"\n\1", s)
# 中文段落分隔:句号/分号后跟新意时换行(保留较短从句)
s = re.sub(r"([。;])\s*", r"\1\n", s)
lines = [ln.strip() for ln in s.splitlines() if ln.strip()]
# 合并过碎的短行(如单独标点)
merged: list[str] = []
for ln in lines:
if merged and len(ln) <= 2 and not re.match(r"^[((]?\d+", ln):
merged[-1] = merged[-1] + ln
else:
merged.append(ln)
return "\n".join(merged)
def _normalize_recommendations(text: str) -> str:
"""建议统一为多行编号列表。"""
if not text:
return ""
s = str(text).strip()
# 已是多行编号
if re.search(r"(?m)^\s*[((]?\d+[\.、))]", s):
return "\n".join(ln.strip() for ln in s.splitlines() if ln.strip())
# 行内编号:1. / 1、 / (1)
items = re.findall(
r"[((]?([1-9]\d?)[\.、))]\s*([^((]*?)(?=(?:[((]?[1-9]\d?[\.、))])|$)",
s,
)
cleaned = [(idx, t.strip(" ;;。 \t")) for idx, t in items if t.strip(" ;;。 \t")]
if len(cleaned) >= 2:
return "\n".join(f"{i}. {t}" for i, (_, t) in enumerate(cleaned, 1))
# 按分号切成条目
chunks = [c.strip(" ;;。") for c in re.split(r"[;;]", s) if c.strip(" ;;。")]
if len(chunks) >= 2:
return "\n".join(f"{i}. {c}" for i, c in enumerate(chunks, 1))
return s
def make_full_report(
study_type: str,
body_part: str,
findings: str,
impression: str,
recommendations: str,
patient_summary: str = "",
) -> str:
"""结构化完整报告:固定四段,便于前端分段渲染。"""
st = STUDY_LABEL.get(study_type.upper(), study_type)
findings_n = _normalize_multiline(findings)
impression_n = _normalize_multiline(impression) or impression
rec_n = _normalize_recommendations(recommendations) or recommendations
header = [
"【影像诊断报告(AI 辅助)】",
f"检查类型:{st}",
f"检查部位:{body_part or '—'}",
]
if patient_summary:
header.append(f"临床摘要:{patient_summary}")
sections = [
"\n".join(header),
"一、影像所见\n" + (findings_n or "—"),
"二、诊断印象\n" + (impression_n or "—"),
"三、建议\n" + (rec_n or "—"),
"四、声明\n本报告由 AI 辅助生成,仅供临床参考,需执业医师审核,不能替代正式报告。",
]
return "\n\n".join(sections)
def generate_imaging_report(req: ImagingReportRequest) -> ImagingReportResponse:
llm = get_llm()
if llm.enabled:
try:
det_lines = []
for d in req.detections:
name = d.label_zh or d.label
box = ",".join(str(int(x)) for x in d.bbox[:4]) if d.bbox else "-"
det_lines.append(f"{name} conf={d.confidence:.0%} box=[{box}]")
data = llm.chat_json(
[
{
"role": "system",
"content": (
"你是三甲医院影像科辅助报告生成器。"
"必须只输出一个 JSON 对象(不要 markdown 代码块),字段:"
"findings(影像所见:多段文字,用换行分隔;先写检查方法与部位,"
"再写病灶描述,再写其余部位阴性所见,勿写成一整段)、"
"impression(诊断印象:1~3 句,可换行)、"
"recommendations(建议:必须用换行的编号列表,如 "
"'1. ...\\n2. ...\\n3. ...',含进一步检查/随访/会诊)、"
"full_report 不要输出(由系统按分段模板拼接)。"
"依据 YOLO 检测结果撰写,专业简洁;"
"明确写明需执业医师审核,不能替代正式报告。"
"禁止编造未提供的患者检验结果。"
),
},
{
"role": "user",
"content": (
f"检查类型={req.study_type}\n"
f"检查部位={req.body_part or '未注明'}\n"
f"患者摘要={req.patient_summary or '无'}\n"
f"规则初诊={req.preliminary_diagnosis}\n"
f"规则所见={req.findings}\n"
f"检测列表:\n" + ("\n".join(det_lines) if det_lines else "(无检出)")
),
},
]
)
findings = _normalize_multiline(str(data.get("findings") or req.findings).strip())
impression = _normalize_multiline(
str(
data.get("impression")
or data.get("preliminary_diagnosis")
or req.preliminary_diagnosis
).strip()
)
rec = _normalize_recommendations(
str(data.get("recommendations") or "建议专科医师复核。").strip()
)
# 始终用分段模板拼完整报告,避免 LLM 输出一整段墙文本
full = make_full_report(
req.study_type, req.body_part, findings, impression, rec, req.patient_summary
)
model_name = llm.info().get("model") or "llm"
logger.info("影像报告已由 LLM 生成 model=%s", model_name)
return ImagingReportResponse(
findings=findings,
impression=impression,
recommendations=rec,
full_report=full,
model_version=f"report-llm:{model_name}",
)
except Exception as e:
logger.warning("影像报告 LLM 失败,回退模板: %s", e)
logger.info("影像报告使用模板模式(LLM 未启用或调用失败) llm_enabled=%s", llm.enabled)
# recommendations 留空,由 imaging API 回退到 build_imaging_texts 的按病灶建议
full = make_full_report(
req.study_type,
req.body_part,
req.findings,
req.preliminary_diagnosis,
"建议结合临床,由影像科/临床医师最终签发。",
req.patient_summary,
)
return ImagingReportResponse(
findings=req.findings,
impression=req.preliminary_diagnosis,
recommendations="",
full_report=full,
model_version="report-template",
)
def generate_decision(req: DecisionRequest) -> DecisionResponse:
rag = get_rag()
query = " ".join(
x for x in [req.diagnosis, req.chief_complaint, req.history, req.imaging_summary] if x
).strip() or "常见病辅助决策"
sources = rag.retrieve(query, top_k=4)
context = rag.build_context(sources)
patient = req.patient
patient_desc = ""
if patient:
patient_desc = f"年龄={patient.age} 性别={patient.gender} 姓名={patient.name or ''}"
llm = get_llm()
if llm.enabled:
try:
data = llm.chat_json(
[
{
"role": "system",
"content": (
"你是临床辅助决策系统。输出严格 JSON,字段:"
"treatment_suggestions(数组,元素含 title,description,confidence),"
"medication_suggestions(数组,元素含 name,dosage,category,confidence),"
"nursing_advice(字符串数组),"
"follow_up_plan(字符串数组),"
"risks(数组,元素含 type,description,level,confidence),"
"conflicts(字符串数组),"
"full_text(字符串)。"
"必须提醒需医师审核;勿编造不存在的检查结果。"
),
},
{
"role": "user",
"content": (
f"患者:{patient_desc}\n"
f"主诉:{req.chief_complaint}\n"
f"病史:{req.history}\n"
f"查体:{req.exam_findings}\n"
f"诊断:{req.diagnosis}\n"
f"用药:{req.medications}\n"
f"影像摘要:{req.imaging_summary}\n"
f"知识库:\n{context or '无'}"
),
},
]
)
return _map_decision(data, sources, engine="fastapi-rag+llm")
except Exception as e:
logger.warning("决策 LLM 失败: %s", e)
return _template_decision(req, sources)
def _map_decision(data: dict[str, Any], sources: list[SourceRef], engine: str) -> DecisionResponse:
risks = []
for r in data.get("risks") or []:
if isinstance(r, dict):
risks.append(
RiskItem(
type=str(r.get("type") or "风险"),
description=str(r.get("description") or ""),
level=str(r.get("level") or "中"),
confidence=float(r.get("confidence") or 0.8),
)
)
return DecisionResponse(
treatment_suggestions=list(data.get("treatment_suggestions") or []),
medication_suggestions=list(data.get("medication_suggestions") or []),
nursing_advice=[str(x) for x in (data.get("nursing_advice") or [])],
follow_up_plan=[str(x) for x in (data.get("follow_up_plan") or [])],
risks=risks,
conflicts=[str(x) for x in (data.get("conflicts") or [])],
sources=sources,
full_text=str(data.get("full_text") or ""),
engine=engine,
)
def _template_decision(req: DecisionRequest, sources: list[SourceRef]) -> DecisionResponse:
dx = req.diagnosis or ""
treatments: list[dict[str, Any]] = []
meds: list[dict[str, Any]] = []
nursing: list[str] = []
follow: list[str] = []
risks: list[RiskItem] = []
if "高血压" in dx:
treatments = [
{"title": "生活方式干预", "description": "低盐饮食,适量有氧运动,控制体重,戒烟限酒", "confidence": 0.95},
{"title": "药物治疗", "description": "可考虑 ACEI/ARB 或 CCB 作为一线方案(需医师确认)", "confidence": 0.9},
]
meds = [
{"name": "氨氯地平", "dosage": "5mg qd", "category": "钙通道阻滞剂", "confidence": 0.9},
{"name": "缬沙坦", "dosage": "80mg qd", "category": "ARB", "confidence": 0.88},
]
nursing = ["监测血压并记录", "宣教服药依从性", "观察头晕、乏力等低血压症状"]
follow = ["1–2 周门诊复查血压", "评估靶器官损害相关检查"]
elif "糖尿病" in dx:
treatments = [
{"title": "饮食运动", "description": "控制总热量与碳水,规律运动", "confidence": 0.95},
{"title": "降糖治疗", "description": "二甲双胍等一线方案需结合肾功能与禁忌", "confidence": 0.9},
]
meds = [{"name": "二甲双胍", "dosage": "0.5g tid", "category": "双胍类", "confidence": 0.92}]
nursing = ["血糖监测指导", "足部护理宣教", "低血糖识别与处理"]
follow = ["2–4 周复诊评估血糖", "定期查 HbA1c"]
elif "肺炎" in dx or "阴影" in dx:
treatments = [
{"title": "抗感染", "description": "根据社区/医院获得性肺炎指南选择抗生素", "confidence": 0.88},
{"title": "支持治疗", "description": "休息、补液、必要时氧疗", "confidence": 0.92},
]
meds = [{"name": "阿莫西林", "dosage": "0.5g tid", "category": "青霉素类", "confidence": 0.85}]
nursing = ["监测体温与呼吸", "叩背排痰指导", "隔离防护宣教(如需要)"]
follow = ["3–5 天评估疗效", "必要时复查胸片"]
elif "结节" in dx:
treatments = [
{"title": "分层随访", "description": "按结节大小与特征选择随访或进一步检查", "confidence": 0.9},
]
nursing = ["戒烟宣教", "避免焦虑,说明随访意义"]
follow = ["3 个月复查 CT", "出现咯血/胸痛及时就诊"]
else:
treatments = [
{"title": "进一步评估", "description": "完善相关检查以明确诊断", "confidence": 0.85},
{"title": "对症处理", "description": "根据症状给予相应支持治疗", "confidence": 0.88},
]
nursing = ["观察病情变化", "用药与生活方式宣教"]
follow = ["按病情 1–2 周复诊", "出现加重症状及时急诊"]
if req.patient and req.patient.age and req.patient.age >= 65:
risks.append(
RiskItem(
type="高龄风险",
description="高龄患者需注意剂量调整、跌倒与多药联用风险",
level="高",
confidence=0.85,
)
)
if "高血压" in dx and req.patient and req.patient.age and req.patient.age > 60:
risks.append(
RiskItem(
type="心血管风险",
description="高血压合并高龄,心血管事件风险增加",
level="中",
confidence=0.8,
)
)
conflicts: list[str] = []
meds_text = req.medications or ""
if "华法林" in meds_text and "阿司匹林" in meds_text:
conflicts.append("警告:华法林与阿司匹林联合使用可能增加出血风险")
if "ACEI" in meds_text and "保钾" in meds_text:
conflicts.append("注意:ACEI 与保钾利尿剂联用可能致高钾血症")
src_hint = ""
if sources:
src_hint = "\n知识库参考:" + ";".join(s.title for s in sources[:3])
full = (
f"诊断相关辅助建议(规则+RAG):{dx or '未明确'}\n"
f"治疗:{'; '.join(t['title'] for t in treatments)}\n"
f"护理:{';'.join(nursing)}\n"
f"随访:{';'.join(follow)}"
f"{src_hint}\n"
"(模板模式,可配置 LLM_API_KEY 启用大模型增强)"
)
return DecisionResponse(
treatment_suggestions=treatments,
medication_suggestions=meds,
nursing_advice=nursing,
follow_up_plan=follow,
risks=risks,
conflicts=conflicts,
sources=sources,
full_text=full,
engine="fastapi-rag-template",
)
+246
View File
@@ -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()
+345
View File
@@ -0,0 +1,345 @@
"""YOLO 权重文件管理:上传、激活、统计可视化数据。"""
from __future__ import annotations
import json
import logging
import re
import shutil
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from app.config import get_settings
logger = logging.getLogger(__name__)
ALLOWED_EXT = {".pt", ".pth", ".onnx", ".engine"}
class YoloManager:
def __init__(self) -> None:
settings = get_settings()
self.root = settings.root_dir
self.weights_dir = (self.root / "data" / "weights").resolve()
self.config_path = (self.root / "data" / "yolo_config.json").resolve()
self.stats_path = (self.root / "data" / "yolo_stats.json").resolve()
self.weights_dir.mkdir(parents=True, exist_ok=True)
self.config_path.parent.mkdir(parents=True, exist_ok=True)
self._config = self._load_json(self.config_path, default={
"active_weight": "",
"demo_mode": "auto", # demo | real | auto
"deleted_weights": [], # 逻辑删除的文件名列表(本地 .pt 仍保留)
"updated_at": None,
})
if not isinstance(self._config.get("deleted_weights"), list):
self._config["deleted_weights"] = []
self._stats = self._load_json(self.stats_path, default=self._empty_stats())
@staticmethod
def _empty_stats() -> dict[str, Any]:
return {
"total_inferences": 0,
"real_count": 0,
"demo_count": 0,
"class_counts": {},
"confidence_buckets": {"0-50": 0, "50-70": 0, "70-85": 0, "85-100": 0},
"recent": [], # last 20 runs
"updated_at": None,
}
@staticmethod
def _load_json(path: Path, default: dict) -> dict:
if not path.is_file():
return dict(default)
try:
data = json.loads(path.read_text(encoding="utf-8"))
if isinstance(data, dict):
merged = dict(default)
merged.update(data)
return merged
except Exception as e:
logger.warning("读取 %s 失败: %s", path, e)
return dict(default)
def _save_config(self) -> None:
self._config["updated_at"] = datetime.now(timezone.utc).isoformat()
self.config_path.write_text(
json.dumps(self._config, ensure_ascii=False, indent=2),
encoding="utf-8",
)
def _save_stats(self) -> None:
self._stats["updated_at"] = datetime.now(timezone.utc).isoformat()
self.stats_path.write_text(
json.dumps(self._stats, ensure_ascii=False, indent=2),
encoding="utf-8",
)
def _deleted_set(self) -> set[str]:
return {str(x) for x in (self._config.get("deleted_weights") or []) if x}
def list_weights(self) -> list[dict[str, Any]]:
"""仅列出未逻辑删除的权重。"""
active = (self._config.get("active_weight") or "").strip()
deleted = self._deleted_set()
items: list[dict[str, Any]] = []
for p in sorted(self.weights_dir.iterdir(), key=lambda x: x.stat().st_mtime, reverse=True):
if not p.is_file() or p.suffix.lower() not in ALLOWED_EXT:
continue
if p.name in deleted:
continue
# 回收目录 / 隐藏文件不展示
if p.name.startswith("."):
continue
st = p.stat()
items.append({
"name": p.name,
"path": str(p),
"size_bytes": st.st_size,
"size_mb": round(st.st_size / (1024 * 1024), 3),
"modified_at": datetime.fromtimestamp(st.st_mtime, tz=timezone.utc).isoformat(),
"active": p.name == active,
"ext": p.suffix.lower(),
"deleted": False,
})
return items
def save_upload(self, filename: str, content: bytes) -> dict[str, Any]:
if not content:
raise ValueError("文件为空")
safe = self._safe_name(filename)
ext = Path(safe).suffix.lower()
if ext not in ALLOWED_EXT:
raise ValueError(f"仅支持权重格式: {', '.join(sorted(ALLOWED_EXT))}")
# 限制 500MB
if len(content) > 500 * 1024 * 1024:
raise ValueError("权重文件不能超过 500MB")
target = self.weights_dir / safe
# 避免覆盖:同名追加时间戳
if target.exists():
stem = target.stem
target = self.weights_dir / f"{stem}_{int(time.time())}{ext}"
safe = target.name
target.write_bytes(content)
logger.info("已保存 YOLO 权重: %s (%d bytes)", target, len(content))
return {
"name": safe,
"path": str(target),
"size_bytes": len(content),
"size_mb": round(len(content) / (1024 * 1024), 3),
"active": False,
}
def activate(self, name: str, demo_mode: str | None = None) -> dict[str, Any]:
path = self.weights_dir / name
if not path.is_file():
raise FileNotFoundError(f"权重不存在: {name}")
if name in self._deleted_set():
raise FileNotFoundError(f"权重已逻辑删除,无法激活: {name}")
if path.suffix.lower() not in ALLOWED_EXT:
raise ValueError("非法权重文件")
self._config["active_weight"] = name
if demo_mode in ("demo", "real", "auto"):
self._config["demo_mode"] = demo_mode
elif not self._config.get("demo_mode"):
self._config["demo_mode"] = "auto"
self._save_config()
return self.status()
def set_demo_mode(self, mode: str) -> dict[str, Any]:
if mode not in ("demo", "real", "auto"):
raise ValueError("demo_mode 仅支持 demo / real / auto")
self._config["demo_mode"] = mode
self._save_config()
return self.status()
def deactivate(self) -> dict[str, Any]:
self._config["active_weight"] = ""
self._save_config()
return self.status()
def delete_weight(self, name: str) -> None:
"""逻辑删除:不删除磁盘文件,仅从可用列表隐藏。"""
path = self.weights_dir / name
if not path.is_file() and name not in self._deleted_set():
raise FileNotFoundError(f"权重不存在: {name}")
if path.is_file() and path.suffix.lower() not in ALLOWED_EXT:
raise ValueError("非法权重文件")
deleted = list(self._config.get("deleted_weights") or [])
if name not in deleted:
deleted.append(name)
self._config["deleted_weights"] = deleted
if (self._config.get("active_weight") or "") == name:
self._config["active_weight"] = ""
self._save_config()
logger.info("YOLO 权重逻辑删除(文件保留): %s path=%s", name, path)
def restore_weight(self, name: str) -> dict[str, Any]:
"""从逻辑删除中恢复(若本地文件仍在)。"""
path = self.weights_dir / name
if not path.is_file():
raise FileNotFoundError(f"本地文件不存在,无法恢复: {name}")
deleted = [x for x in (self._config.get("deleted_weights") or []) if x != name]
self._config["deleted_weights"] = deleted
self._save_config()
logger.info("YOLO 权重已从逻辑删除恢复: %s", name)
return self.status()
def active_weight_path(self) -> Path | None:
name = (self._config.get("active_weight") or "").strip()
if not name:
# 兼容环境变量
settings = get_settings()
return settings.yolo_weights_path
if name in self._deleted_set():
return None
path = self.weights_dir / name
return path if path.is_file() else None
def effective_demo_mode(self) -> str:
mode = (self._config.get("demo_mode") or "auto").strip().lower()
if mode in ("demo", "real", "auto"):
return mode
return get_settings().demo_mode
def status(self, detector_info: dict[str, Any] | None = None) -> dict[str, Any]:
active_path = self.active_weight_path()
weights = self.list_weights()
body: dict[str, Any] = {
"weights_dir": str(self.weights_dir),
"active_weight": self._config.get("active_weight") or "",
"active_path": str(active_path) if active_path else None,
"active_exists": active_path is not None and active_path.is_file(),
"demo_mode": self.effective_demo_mode(),
"weights_count": len(weights),
"weights": weights,
"updated_at": self._config.get("updated_at"),
}
if detector_info:
body.update(detector_info)
return body
def record_inference(
self,
mode: str,
detections: list[Any],
study_type: str = "",
model_version: str = "",
) -> None:
self._stats["total_inferences"] = int(self._stats.get("total_inferences") or 0) + 1
if mode == "real":
self._stats["real_count"] = int(self._stats.get("real_count") or 0) + 1
else:
self._stats["demo_count"] = int(self._stats.get("demo_count") or 0) + 1
class_counts: dict[str, int] = self._stats.setdefault("class_counts", {})
buckets: dict[str, int] = self._stats.setdefault(
"confidence_buckets",
{"0-50": 0, "50-70": 0, "70-85": 0, "85-100": 0},
)
labels: list[str] = []
confs: list[float] = []
for d in detections or []:
if hasattr(d, "label"):
label = getattr(d, "label_zh", None) or d.label
conf = float(d.confidence)
elif isinstance(d, dict):
label = d.get("label_zh") or d.get("label") or "unknown"
conf = float(d.get("confidence") or 0)
else:
continue
labels.append(str(label))
confs.append(conf)
class_counts[str(label)] = int(class_counts.get(str(label), 0)) + 1
pct = conf * 100
if pct < 50:
buckets["0-50"] = buckets.get("0-50", 0) + 1
elif pct < 70:
buckets["50-70"] = buckets.get("50-70", 0) + 1
elif pct < 85:
buckets["70-85"] = buckets.get("70-85", 0) + 1
else:
buckets["85-100"] = buckets.get("85-100", 0) + 1
recent = self._stats.setdefault("recent", [])
recent.insert(0, {
"time": datetime.now(timezone.utc).isoformat(),
"mode": mode,
"study_type": study_type,
"model_version": model_version,
"detection_count": len(labels),
"labels": labels[:10],
"avg_confidence": round(sum(confs) / len(confs), 4) if confs else 0,
"weight": self._config.get("active_weight") or "",
})
self._stats["recent"] = recent[:30]
self._save_stats()
def visualization(self) -> dict[str, Any]:
"""供前端 ECharts 使用的聚合数据。"""
class_counts = self._stats.get("class_counts") or {}
buckets = self._stats.get("confidence_buckets") or {}
real = int(self._stats.get("real_count") or 0)
demo = int(self._stats.get("demo_count") or 0)
total = int(self._stats.get("total_inferences") or 0)
class_pie = [
{"name": k, "value": v}
for k, v in sorted(class_counts.items(), key=lambda x: -x[1])
]
conf_bar = [
{"name": k, "value": int(buckets.get(k, 0))}
for k in ["0-50", "50-70", "70-85", "85-100"]
]
mode_pie = [
{"name": "真实权重推理", "value": real},
{"name": "演示模式", "value": demo},
]
# 近 10 次检测数折线
recent = list(reversed(self._stats.get("recent") or []))[-15:]
trend = {
"times": [ (r.get("time") or "")[11:19] for r in recent ],
"counts": [ int(r.get("detection_count") or 0) for r in recent ],
"modes": [ r.get("mode") or "" for r in recent ],
}
active = self.active_weight_path()
return {
"summary": {
"total_inferences": total,
"real_count": real,
"demo_count": demo,
"real_ratio": round(real / total, 4) if total else 0,
"active_weight": self._config.get("active_weight") or "",
"demo_mode": self.effective_demo_mode(),
"active_exists": bool(active and active.is_file()),
"weights_count": len(self.list_weights()),
},
"class_distribution": class_pie,
"confidence_distribution": conf_bar,
"mode_distribution": mode_pie,
"inference_trend": trend,
"recent": self._stats.get("recent") or [],
"updated_at": self._stats.get("updated_at"),
}
def reset_stats(self) -> None:
self._stats = self._empty_stats()
self._save_stats()
@staticmethod
def _safe_name(filename: str) -> str:
name = Path(filename or "weights.pt").name
name = re.sub(r"[^\w.\-()+]", "_", name)
if not name or name in (".", ".."):
name = f"weights_{int(time.time())}.pt"
return name
_manager: YoloManager | None = None
def get_yolo_manager() -> YoloManager:
global _manager
if _manager is None:
_manager = YoloManager()
return _manager