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 @@
"""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