26-7-31-1
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""AI service implementations."""
|
||||
@@ -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
|
||||
@@ -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},
|
||||
}
|
||||
@@ -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
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user