346 lines
13 KiB
Python
346 lines
13 KiB
Python
"""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
|