"""实时摄像头推理:支持 YOLO / Faster R-CNN 权重。""" 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 ROOT = Path(__file__).resolve().parent COLOR_RESOLUTIONS = [(1280, 720), (848, 480), (640, 480), (1920, 1080)] 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_exp" / "weights" / "best.pt", ROOT / "ultralytics" / "runs" / "detect" / "runs" / "yolov8_exp" / "weights" / "last.pt", ROOT / "runs" / "detect" / "yolov8_exp" / "weights" / "best.pt", ROOT / "yolov8n.pt", ] for path in candidates: if path.exists(): return str(path) return str(ROOT / "yolov8n.pt") def list_realsense() -> list[tuple[str, str]]: try: import pyrealsense2 as rs except Exception: return [] items = [] for dev in rs.context().query_devices(): name = dev.get_info(rs.camera_info.name) serial = dev.get_info(rs.camera_info.serial_number) items.append((f"{name} [{serial}]", serial)) return items class RealSenseRGB: """与 capture.py 一致:同时开左右红外 + RGB,界面只用 RGB。""" def __init__(self, width: int, height: int, fps: int = 30, serial: str | None = None) -> None: import pyrealsense2 as rs self.rs = rs self.pipeline = None self.size = (width, height) last_error: Exception | None = None attempts = [(width, height, fps)] + [(w, h, fps) for w, h in COLOR_RESOLUTIONS if (w, h) != (width, height)] # 推理只要 RGB:先试仅彩色(快、省带宽);失败再试红外+RGB(与 capture 相同) modes = ("rgb_only", "rgb_ir") for mode in modes: for w, h, f in attempts: pipeline = rs.pipeline() config = rs.config() if serial: config.enable_device(serial) try: if mode == "rgb_ir": config.enable_stream(rs.stream.infrared, 1, w, h, rs.format.y8, f) config.enable_stream(rs.stream.infrared, 2, w, h, rs.format.y8, f) cw, ch = (w, h) if (w, h) in COLOR_RESOLUTIONS else (1280, 720) config.enable_stream(rs.stream.color, cw, ch, rs.format.bgr8, f) profile = pipeline.start(config) try: depth_sensor = profile.get_device().first_depth_sensor() if depth_sensor.supports(rs.option.emitter_enabled): depth_sensor.set_option(rs.option.emitter_enabled, 0.0) except Exception: pass for _ in range(5): pipeline.wait_for_frames(1500) self.pipeline = pipeline self.size = (cw, ch) return except Exception as exc: last_error = exc try: pipeline.stop() except Exception: pass raise RuntimeError( f"无法打开 RealSense RGB(请先关掉 capture.py / RealSense Viewer): {last_error}" ) def read(self) -> tuple[bool, np.ndarray | None]: if self.pipeline is None: return False, None try: frames = self.pipeline.wait_for_frames(1000) except Exception: return False, None color = frames.get_color_frame() if not color: return False, None return True, np.asanyarray(color.get_data()) def release(self) -> None: if self.pipeline is None: return try: self.pipeline.stop() except Exception: pass self.pipeline = None class Webcam: def __init__(self, index: int, width: int, height: int, fps: int = 30) -> None: cap = None for backend in (cv2.CAP_DSHOW, cv2.CAP_MSMF, cv2.CAP_ANY): trial = cv2.VideoCapture(index, backend) if trial.isOpened(): cap = trial break trial.release() if cap is None or not cap.isOpened(): raise RuntimeError(f"无法打开摄像头 {index}。本机可试编号 1(0 常被 RealSense UVC 占住)") cap.set(cv2.CAP_PROP_FRAME_WIDTH, width) cap.set(cv2.CAP_PROP_FRAME_HEIGHT, height) cap.set(cv2.CAP_PROP_FPS, fps) ok, frame = False, None for _ in range(15): ok, frame = cap.read() if ok and frame is not None: break time.sleep(0.05) if not ok or frame is None: cap.release() raise RuntimeError(f"摄像头 {index} 已打开但读不到画面,请换编号或关掉占用程序") self.cap = cap def read(self) -> tuple[bool, np.ndarray | None]: ok, frame = self.cap.read() return ok, frame if ok else None def release(self) -> None: self.cap.release() def open_source(source: str, width: int, height: int, serial: str | None, cam_index: int): if source == "realsense": return RealSenseRGB(width, height, 30, serial) return Webcam(cam_index, width, height, 30) class InferApp(tk.Tk): def __init__(self) -> None: super().__init__() self.title("实时检测推理(YOLO / Faster R-CNN)") self.geometry("1180x820") self.minsize(980, 700) self.model = None 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.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_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 "") # 本机普通摄像头能出图的多半是 1 self.var_camera = tk.IntVar(value=1) self.var_width = tk.IntVar(value=1280) self.var_height = tk.IntVar(value=720) self.var_conf = tk.DoubleVar(value=0.25) self.var_iou = tk.DoubleVar(value=0.7) self.var_imgsz = tk.IntVar(value=640) 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) 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)) src = ttk.Frame(panel) src.grid(row=2, 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="RealSense").grid(row=3, 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=3, column=1, columnspan=3, sticky="ew", pady=(8, 0)) ttk.Button(panel, text="刷新设备", command=self.refresh_devices).grid(row=3, column=4, sticky="w", pady=(8, 0)) ttk.Label(panel, text="摄像头编号").grid(row=3, column=5, sticky="w", pady=(8, 0)) ttk.Spinbox(panel, from_=0, to=8, textvariable=self.var_camera, width=6).grid( row=3, column=6, sticky="w", pady=(8, 0) ) ttk.Label(panel, text="分辨率").grid(row=4, column=0, sticky="w", pady=(8, 0)) ttk.Entry(panel, textvariable=self.var_width, width=8).grid(row=4, column=1, sticky="w", pady=(8, 0)) ttk.Label(panel, text="x").grid(row=4, column=2, pady=(8, 0)) ttk.Entry(panel, textvariable=self.var_height, width=8).grid(row=4, column=3, sticky="w", pady=(8, 0)) ttk.Label(panel, text="imgsz").grid(row=4, column=4, sticky="w", pady=(8, 0)) ttk.Entry(panel, textvariable=self.var_imgsz, width=8).grid(row=4, column=5, sticky="w", pady=(8, 0)) ttk.Label(panel, text="设备").grid(row=4, column=6, sticky="w", pady=(8, 0)) ttk.Combobox(panel, textvariable=self.var_device, values=["0", "cpu"], width=8, state="readonly").grid( row=4, column=7, sticky="w", pady=(8, 0) ) ttk.Label(panel, text="置信度").grid(row=5, column=0, sticky="w", pady=(8, 0)) ttk.Scale(panel, from_=0.05, to=0.95, variable=self.var_conf, orient="horizontal").grid( row=5, column=1, columnspan=2, sticky="ew", pady=(8, 0) ) self.lbl_conf = ttk.Label(panel, text="0.25") self.lbl_conf.grid(row=5, column=3, sticky="w", pady=(8, 0)) self.var_conf.trace_add("write", lambda *_: self.lbl_conf.config(text=f"{float(self.var_conf.get()):.2f}")) ttk.Label(panel, text="IoU").grid(row=5, column=4, sticky="w", pady=(8, 0)) ttk.Entry(panel, textvariable=self.var_iou, width=8).grid(row=5, column=5, 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 if backend != "auto" else "yolo")) 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 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") use_serial = serial or devices[0][1] cam = open_source("realsense", width, height, use_serial, 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") w, h = getattr(cam, "size", (self.var_width.get(), self.var_height.get())) self.var_status.set(f"摄像头已打开 {w}x{h}(预览中)。确认画面后点「开始推理」") 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.btn_infer.config(state="disabled") self.var_status.set(f"正在加载模型: {model_path.name} …") # 取一帧热身,避免和采集线程同时 wait_for_frames 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.var_conf.get()), ) if warmup is not None: model.plot( warmup, conf=float(self.var_conf.get()), iou=float(self.var_iou.get()), imgsz=int(self.var_imgsz.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 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 plotted = frame det = 0 if self.infer_enabled and self.model is not None: try: dets = self.model.detect( frame, conf=float(self.var_conf.get()), iou=float(self.var_iou.get()), imgsz=int(self.var_imgsz.get()), ) plotted = frame.copy() for d in dets: x1, y1, x2, y2 = d.xyxy cv2.rectangle(plotted, (x1, y1), (x2, y2), (0, 255, 0), 2) cv2.putText( plotted, f"{d.conf:.2f}", (x1, max(20, y1 - 6)), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2, ) det = len(dets) except Exception as exc: self.after(0, lambda e=str(exc): self.var_status.set(f"推理错误: {e}")) time.sleep(0.05) continue n += 1 if n >= 8: now = time.perf_counter() self.fps = n / max(1e-6, now - t0) t0, n = now, 0 show = plotted.copy() mode = "DET" if self.infer_enabled else "PREVIEW" cv2.putText( show, f"{mode} FPS {self.fps:.1f} det={det} conf>={float(self.var_conf.get()):.2f}", (16, 36), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 255, 0) if self.infer_enabled else (0, 200, 255), 2, ) with self.lock: self.latest = show 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_realtime" out_dir.mkdir(parents=True, exist_ok=True) out = out_dir / f"shot_{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__": InferApp().mainloop()