26-7-31-1
This commit is contained in:
@@ -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