480 lines
19 KiB
Python
480 lines
19 KiB
Python
"""工件螺纹孔二阶段视频推理:YOLO / Faster R-CNN 定位 + HoughCircles。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import threading
|
|
import time
|
|
import tkinter as tk
|
|
from pathlib import Path
|
|
from tkinter import filedialog, messagebox, ttk
|
|
|
|
import cv2
|
|
import numpy as np
|
|
from PIL import Image, ImageTk
|
|
|
|
from detectors import guess_backend, load_detector
|
|
from hole_detect import HoleStabilizer, detect_holes, load_config
|
|
from infer_realtime import (
|
|
list_realsense,
|
|
open_source,
|
|
)
|
|
|
|
ROOT = Path(__file__).resolve().parent
|
|
|
|
|
|
def default_model_path(backend: str = "yolo") -> str:
|
|
if backend == "faster_rcnn":
|
|
candidates = [
|
|
ROOT / "runs" / "faster_rcnn" / "gongjian_frcnn" / "weights" / "best.pth",
|
|
ROOT / "runs" / "faster_rcnn" / "gongjian_frcnn" / "weights" / "last.pth",
|
|
]
|
|
for path in candidates:
|
|
if path.exists():
|
|
return str(path)
|
|
return str(ROOT / "runs" / "faster_rcnn" / "gongjian_frcnn" / "weights" / "best.pth")
|
|
candidates = [
|
|
ROOT / "ultralytics" / "runs" / "detect" / "runs" / "yolov8_gongjian" / "weights" / "best.pt",
|
|
ROOT / "ultralytics" / "runs" / "detect" / "runs" / "yolov8_gongjian" / "weights" / "last.pt",
|
|
ROOT / "ultralytics" / "runs" / "detect" / "runs" / "yolov8_exp" / "weights" / "best.pt",
|
|
ROOT / "yolov8n.pt",
|
|
]
|
|
for path in candidates:
|
|
if path.exists():
|
|
return str(path)
|
|
return str(ROOT / "yolov8n.pt")
|
|
|
|
|
|
def default_config_path() -> str:
|
|
path = ROOT / "hole_detect_config.yaml"
|
|
return str(path)
|
|
|
|
|
|
class HoleInferApp(tk.Tk):
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.title("工件螺纹孔二阶段检测")
|
|
self.geometry("1180x860")
|
|
self.minsize(980, 720)
|
|
|
|
self.model = None
|
|
self.cfg: dict = {}
|
|
self.cam = None
|
|
self.running = False
|
|
self.opening = False
|
|
self.infer_enabled = False
|
|
self.photo: ImageTk.PhotoImage | None = None
|
|
self.latest: np.ndarray | None = None
|
|
self.last_result: dict | None = None
|
|
self.stabilizer = HoleStabilizer()
|
|
self.fps = 0.0
|
|
self.lock = threading.Lock()
|
|
self.cam_lock = threading.Lock()
|
|
self.fail_reads = 0
|
|
self.rs_devices = list_realsense()
|
|
|
|
self.var_backend = tk.StringVar(value="yolo")
|
|
self.var_model = tk.StringVar(value=default_model_path("yolo"))
|
|
self.var_num_classes = tk.IntVar(value=2)
|
|
self.var_config = tk.StringVar(value=default_config_path())
|
|
self.var_source = tk.StringVar(value="realsense" if self.rs_devices else "webcam")
|
|
self.var_rs = tk.StringVar(value=self.rs_devices[0][0] if self.rs_devices else "")
|
|
self.var_camera = tk.IntVar(value=1)
|
|
self.var_width = tk.IntVar(value=1280)
|
|
self.var_height = tk.IntVar(value=720)
|
|
self.var_device = tk.StringVar(value="0")
|
|
self.var_status = tk.StringVar(value="就绪:打开摄像头 → 开始推理")
|
|
|
|
self._build()
|
|
self.protocol("WM_DELETE_WINDOW", self.on_close)
|
|
self.after(33, self.refresh_view)
|
|
self.reload_config(silent=True)
|
|
|
|
def _build(self) -> None:
|
|
self.columnconfigure(0, weight=1)
|
|
self.rowconfigure(0, weight=1)
|
|
|
|
view = ttk.Frame(self, padding=8)
|
|
view.grid(row=0, column=0, sticky="nsew")
|
|
view.columnconfigure(0, weight=1)
|
|
view.rowconfigure(0, weight=1)
|
|
self.canvas = tk.Label(view, bg="#111", fg="#eee", text="检测画面", width=80, height=20)
|
|
self.canvas.grid(row=0, column=0, sticky="nsew")
|
|
ttk.Label(view, textvariable=self.var_status).grid(row=1, column=0, sticky="w", pady=(8, 0))
|
|
|
|
panel = ttk.LabelFrame(self, text="二阶段检测参数", padding=10)
|
|
panel.grid(row=1, column=0, sticky="ew", padx=8, pady=(0, 8))
|
|
for c in range(8):
|
|
panel.columnconfigure(c, weight=1)
|
|
|
|
ttk.Label(panel, text="算法").grid(row=0, column=0, sticky="w")
|
|
ttk.Combobox(
|
|
panel,
|
|
textvariable=self.var_backend,
|
|
values=["yolo", "faster_rcnn", "auto"],
|
|
width=12,
|
|
state="readonly",
|
|
).grid(row=0, column=1, sticky="w")
|
|
ttk.Label(panel, text="RCNN类数(+背景)").grid(row=0, column=2, sticky="w")
|
|
ttk.Entry(panel, textvariable=self.var_num_classes, width=6).grid(row=0, column=3, sticky="w")
|
|
self.var_backend.trace_add("write", lambda *_: self.on_backend_change())
|
|
|
|
ttk.Label(panel, text="检测权重").grid(row=1, column=0, sticky="w", pady=(8, 0))
|
|
ttk.Entry(panel, textvariable=self.var_model).grid(row=1, column=1, columnspan=5, sticky="ew", padx=4, pady=(8, 0))
|
|
ttk.Button(panel, text="浏览…", command=self.browse_model).grid(row=1, column=6, sticky="ew", pady=(8, 0))
|
|
ttk.Button(panel, text="自动查找", command=self.autofill_model).grid(row=1, column=7, sticky="ew", padx=(4, 0), pady=(8, 0))
|
|
|
|
ttk.Label(panel, text="检测配置").grid(row=2, column=0, sticky="w", pady=(8, 0))
|
|
ttk.Entry(panel, textvariable=self.var_config).grid(row=2, column=1, columnspan=5, sticky="ew", padx=4, pady=(8, 0))
|
|
ttk.Button(panel, text="浏览…", command=self.browse_config).grid(row=2, column=6, sticky="ew", pady=(8, 0))
|
|
ttk.Button(panel, text="重新加载", command=self.reload_config).grid(row=2, column=7, sticky="ew", padx=(4, 0), pady=(8, 0))
|
|
|
|
ttk.Label(panel, text="视频源").grid(row=3, column=0, sticky="w", pady=(8, 0))
|
|
src = ttk.Frame(panel)
|
|
src.grid(row=3, column=1, columnspan=3, sticky="w", pady=(8, 0))
|
|
ttk.Radiobutton(src, text="RealSense RGB", variable=self.var_source, value="realsense").pack(side="left")
|
|
ttk.Radiobutton(src, text="普通摄像头", variable=self.var_source, value="webcam").pack(side="left", padx=(10, 0))
|
|
ttk.Label(panel, text="设备").grid(row=3, column=4, sticky="w", pady=(8, 0))
|
|
ttk.Combobox(panel, textvariable=self.var_device, values=["0", "cpu"], width=8, state="readonly").grid(
|
|
row=3, column=5, sticky="w", pady=(8, 0)
|
|
)
|
|
|
|
ttk.Label(panel, text="RealSense").grid(row=4, column=0, sticky="w", pady=(8, 0))
|
|
self.cmb_rs = ttk.Combobox(
|
|
panel,
|
|
textvariable=self.var_rs,
|
|
values=[x[0] for x in self.rs_devices],
|
|
state="readonly",
|
|
)
|
|
self.cmb_rs.grid(row=4, column=1, columnspan=3, sticky="ew", pady=(8, 0))
|
|
ttk.Button(panel, text="刷新设备", command=self.refresh_devices).grid(row=4, column=4, sticky="w", pady=(8, 0))
|
|
ttk.Label(panel, text="摄像头编号").grid(row=4, column=5, sticky="w", pady=(8, 0))
|
|
ttk.Spinbox(panel, from_=0, to=8, textvariable=self.var_camera, width=6).grid(
|
|
row=4, column=6, sticky="w", pady=(8, 0)
|
|
)
|
|
|
|
ttk.Label(panel, text="分辨率").grid(row=5, column=0, sticky="w", pady=(8, 0))
|
|
ttk.Entry(panel, textvariable=self.var_width, width=8).grid(row=5, column=1, sticky="w", pady=(8, 0))
|
|
ttk.Label(panel, text="x").grid(row=5, column=2, pady=(8, 0))
|
|
ttk.Entry(panel, textvariable=self.var_height, width=8).grid(row=5, column=3, sticky="w", pady=(8, 0))
|
|
|
|
btns = ttk.Frame(panel)
|
|
btns.grid(row=6, column=0, columnspan=8, sticky="ew", pady=(12, 0))
|
|
self.btn_open = ttk.Button(btns, text="打开摄像头", command=self.open_camera)
|
|
self.btn_open.pack(side="left")
|
|
self.btn_infer = ttk.Button(btns, text="开始推理", command=self.start_infer, state="disabled")
|
|
self.btn_infer.pack(side="left", padx=8)
|
|
self.btn_stop = ttk.Button(btns, text="停止", command=self.stop, state="disabled")
|
|
self.btn_stop.pack(side="left")
|
|
ttk.Button(btns, text="保存当前帧", command=self.save_frame).pack(side="left", padx=8)
|
|
|
|
def refresh_devices(self) -> None:
|
|
self.rs_devices = list_realsense()
|
|
labels = [x[0] for x in self.rs_devices]
|
|
self.cmb_rs["values"] = labels
|
|
if labels:
|
|
self.var_rs.set(labels[0])
|
|
self.var_status.set(f"找到 {len(labels)} 台 RealSense")
|
|
else:
|
|
self.var_rs.set("")
|
|
self.var_status.set("未检测到 RealSense")
|
|
|
|
def current_rs_serial(self) -> str | None:
|
|
label = self.var_rs.get()
|
|
for text, serial in self.rs_devices:
|
|
if text == label:
|
|
return serial
|
|
return None
|
|
|
|
def on_backend_change(self) -> None:
|
|
backend = self.var_backend.get()
|
|
if backend == "auto":
|
|
return
|
|
self.var_model.set(default_model_path(backend))
|
|
|
|
def browse_model(self) -> None:
|
|
initial = self.var_model.get()
|
|
initial_dir = str(Path(initial).parent) if initial else str(ROOT)
|
|
path = filedialog.askopenfilename(
|
|
title="选择检测权重",
|
|
initialdir=initial_dir,
|
|
filetypes=[("PyTorch 权重", "*.pt *.pth"), ("所有文件", "*.*")],
|
|
)
|
|
if path:
|
|
self.var_model.set(path)
|
|
if self.var_backend.get() == "auto":
|
|
self.var_status.set(f"auto 将按文件识别为: {guess_backend(path)}")
|
|
|
|
def autofill_model(self) -> None:
|
|
backend = self.var_backend.get()
|
|
if backend == "auto":
|
|
backend = "yolo"
|
|
path = default_model_path(backend)
|
|
self.var_model.set(path)
|
|
self.var_status.set(f"已填入权重: {path}")
|
|
|
|
def browse_config(self) -> None:
|
|
initial = self.var_config.get()
|
|
initial_dir = str(Path(initial).parent) if initial else str(ROOT)
|
|
path = filedialog.askopenfilename(
|
|
title="选择检测配置",
|
|
initialdir=initial_dir,
|
|
filetypes=[("YAML", "*.yaml;*.yml"), ("所有文件", "*.*")],
|
|
)
|
|
if path:
|
|
self.var_config.set(path)
|
|
self.reload_config()
|
|
|
|
def reload_config(self, silent: bool = False) -> None:
|
|
try:
|
|
path = Path(self.var_config.get().strip())
|
|
if not path.is_absolute():
|
|
path = ROOT / path
|
|
self.cfg = load_config(path)
|
|
self.var_config.set(str(path))
|
|
self.stabilizer = HoleStabilizer(self.cfg)
|
|
if not silent:
|
|
self.var_status.set(f"已加载配置: {path.name}")
|
|
except Exception as exc:
|
|
self.cfg = {}
|
|
self.stabilizer = HoleStabilizer()
|
|
if not silent:
|
|
messagebox.showerror("配置", str(exc))
|
|
self.var_status.set(f"配置加载失败: {exc}")
|
|
|
|
def release_cam(self) -> None:
|
|
with self.cam_lock:
|
|
if self.cam is not None:
|
|
try:
|
|
self.cam.release()
|
|
except Exception:
|
|
pass
|
|
self.cam = None
|
|
|
|
def open_camera(self) -> None:
|
|
if self.running or self.opening:
|
|
return
|
|
self.opening = True
|
|
self.btn_open.config(state="disabled")
|
|
self.var_status.set("正在打开摄像头…")
|
|
self.update_idletasks()
|
|
|
|
source = self.var_source.get()
|
|
width = int(self.var_width.get())
|
|
height = int(self.var_height.get())
|
|
serial = self.current_rs_serial()
|
|
cam_index = int(self.var_camera.get())
|
|
|
|
def worker() -> None:
|
|
cam = None
|
|
err = ""
|
|
try:
|
|
self.release_cam()
|
|
if source == "realsense":
|
|
devices = list_realsense()
|
|
if not devices:
|
|
raise RuntimeError("没有 RealSense,可改选普通摄像头编号 1")
|
|
cam = open_source("realsense", width, height, serial or devices[0][1], cam_index)
|
|
else:
|
|
cam = open_source("webcam", width, height, None, cam_index)
|
|
ok, frame = cam.read()
|
|
if not ok or frame is None:
|
|
raise RuntimeError("摄像头打开后读不到画面")
|
|
except Exception as exc:
|
|
if cam is not None:
|
|
try:
|
|
cam.release()
|
|
except Exception:
|
|
pass
|
|
cam = None
|
|
err = str(exc)
|
|
self.after(0, lambda: self._open_done(cam, err))
|
|
|
|
threading.Thread(target=worker, daemon=True).start()
|
|
|
|
def _open_done(self, cam, err: str) -> None:
|
|
self.opening = False
|
|
if err or cam is None:
|
|
self.btn_open.config(state="normal")
|
|
self.btn_infer.config(state="disabled")
|
|
self.btn_stop.config(state="disabled")
|
|
self.var_status.set(err or "打开失败")
|
|
messagebox.showerror("摄像头", err or "打开失败")
|
|
return
|
|
with self.cam_lock:
|
|
self.cam = cam
|
|
self.running = True
|
|
self.infer_enabled = False
|
|
self.fail_reads = 0
|
|
self.btn_open.config(state="disabled")
|
|
self.btn_infer.config(state="normal")
|
|
self.btn_stop.config(state="normal")
|
|
self.var_status.set("摄像头已打开。点「开始推理」加载权重并检测螺纹孔")
|
|
threading.Thread(target=self.loop_capture, daemon=True).start()
|
|
|
|
def start_infer(self) -> None:
|
|
if not self.running or self.cam is None:
|
|
messagebox.showwarning("推理", "请先打开摄像头")
|
|
return
|
|
model_path = Path(self.var_model.get().strip())
|
|
if not model_path.is_absolute():
|
|
model_path = ROOT / model_path
|
|
if not model_path.exists():
|
|
messagebox.showerror("权重", f"找不到权重:\n{model_path}")
|
|
return
|
|
self.reload_config(silent=True)
|
|
if not self.cfg:
|
|
messagebox.showerror("配置", "请先加载有效的 hole_detect_config.yaml")
|
|
return
|
|
|
|
self.btn_infer.config(state="disabled")
|
|
self.var_status.set(f"正在加载模型: {model_path.name} …")
|
|
warmup = None
|
|
with self.lock:
|
|
if self.latest is not None:
|
|
warmup = self.latest.copy()
|
|
|
|
def load() -> None:
|
|
try:
|
|
backend = self.var_backend.get()
|
|
model = load_detector(
|
|
backend=backend,
|
|
weights=model_path,
|
|
device=self.var_device.get(),
|
|
num_classes=int(self.var_num_classes.get()),
|
|
conf=float(self.cfg.get("conf", 0.25)),
|
|
)
|
|
if warmup is not None:
|
|
detect_holes(warmup, model, self.cfg, device=self.var_device.get())
|
|
except Exception as exc:
|
|
self.after(0, lambda e=str(exc): self._infer_load_failed(e))
|
|
return
|
|
self.model = model
|
|
self.infer_enabled = True
|
|
self.stabilizer.reset()
|
|
used = getattr(model, "backend", backend)
|
|
self.after(0, lambda: self.var_status.set(f"推理中 backend={used} 模型={model_path.name}"))
|
|
|
|
threading.Thread(target=load, daemon=True).start()
|
|
|
|
def _infer_load_failed(self, msg: str) -> None:
|
|
self.infer_enabled = False
|
|
self.model = None
|
|
self.btn_infer.config(state="normal")
|
|
self.var_status.set(msg)
|
|
messagebox.showerror("模型加载失败", msg)
|
|
|
|
def stop(self) -> None:
|
|
self.running = False
|
|
self.infer_enabled = False
|
|
self.var_status.set("正在停止…")
|
|
|
|
def loop_capture(self) -> None:
|
|
t0, n = time.perf_counter(), 0
|
|
while self.running:
|
|
with self.cam_lock:
|
|
cam = self.cam
|
|
if cam is None:
|
|
break
|
|
ok, frame = cam.read()
|
|
if not ok or frame is None:
|
|
self.fail_reads += 1
|
|
if self.fail_reads >= 30:
|
|
self.after(0, lambda: self.var_status.set("连续读帧失败,请停止后重开"))
|
|
self.fail_reads = 0
|
|
time.sleep(0.02)
|
|
continue
|
|
self.fail_reads = 0
|
|
|
|
show = frame
|
|
result = None
|
|
if self.infer_enabled and self.model is not None:
|
|
raw = detect_holes(frame, self.model, self.cfg, device=self.var_device.get())
|
|
result = self.stabilizer.update(frame, raw)
|
|
show = result["draw_img"]
|
|
self.last_result = result
|
|
|
|
n += 1
|
|
if n >= 8:
|
|
now = time.perf_counter()
|
|
self.fps = n / max(1e-6, now - t0)
|
|
t0, n = now, 0
|
|
|
|
overlay = show.copy()
|
|
if result is None:
|
|
mode = "PREVIEW"
|
|
color = (0, 200, 255)
|
|
info = ""
|
|
elif result["status"]:
|
|
mode = "OK"
|
|
color = (0, 255, 0)
|
|
pts = result["hole_points"]
|
|
info = f" holes={pts}"
|
|
else:
|
|
mode = "FAIL"
|
|
color = (0, 0, 255)
|
|
info = f" {result['msg']}"
|
|
cv2.putText(
|
|
overlay,
|
|
f"{mode} FPS {self.fps:.1f}{info}",
|
|
(16, overlay.shape[0] - 20),
|
|
cv2.FONT_HERSHEY_SIMPLEX,
|
|
0.7,
|
|
color,
|
|
2,
|
|
)
|
|
with self.lock:
|
|
self.latest = overlay
|
|
|
|
self.release_cam()
|
|
self.model = None
|
|
self.infer_enabled = False
|
|
self.after(0, self._stopped_ui)
|
|
|
|
def _stopped_ui(self) -> None:
|
|
self.btn_open.config(state="normal")
|
|
self.btn_infer.config(state="disabled")
|
|
self.btn_stop.config(state="disabled")
|
|
self.var_status.set("已停止")
|
|
|
|
def refresh_view(self) -> None:
|
|
frame = None
|
|
with self.lock:
|
|
if self.latest is not None:
|
|
frame = self.latest.copy()
|
|
if frame is not None:
|
|
try:
|
|
rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
|
h, w = rgb.shape[:2]
|
|
max_w = max(int(self.canvas.winfo_width()), 640)
|
|
max_h = max(int(self.canvas.winfo_height()), 360)
|
|
scale = min(max_w / w, max_h / h, 1.0)
|
|
size = (max(1, int(w * scale)), max(1, int(h * scale)))
|
|
img = Image.fromarray(rgb).resize(size, Image.Resampling.BILINEAR)
|
|
self.photo = ImageTk.PhotoImage(img)
|
|
self.canvas.config(image=self.photo, text="")
|
|
except Exception as exc:
|
|
self.var_status.set(f"画面刷新失败: {exc}")
|
|
self.after(33, self.refresh_view)
|
|
|
|
def save_frame(self) -> None:
|
|
with self.lock:
|
|
frame = None if self.latest is None else self.latest.copy()
|
|
if frame is None:
|
|
messagebox.showwarning("保存", "当前没有可保存的画面")
|
|
return
|
|
out_dir = ROOT / "runs" / "predict_holes"
|
|
out_dir.mkdir(parents=True, exist_ok=True)
|
|
out = out_dir / f"hole_{int(time.time())}.jpg"
|
|
cv2.imwrite(str(out), frame)
|
|
self.var_status.set(f"已保存: {out}")
|
|
|
|
def on_close(self) -> None:
|
|
self.running = False
|
|
self.infer_enabled = False
|
|
self.release_cam()
|
|
self.destroy()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
HoleInferApp().mainloop()
|