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