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