From aa96ed1bbec4d578d4c737e468b762b1e83d6c48 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Sun, 21 Dec 2025 21:44:09 +0800 Subject: [PATCH 1/7] Enhance config.py with environment variable support Refactor configuration management to support environment variables and improve readability. --- .../config.py | 99 +++++++++++++++---- 1 file changed, 79 insertions(+), 20 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/config.py b/src/Driving_Accident_Video_Recognition/config.py index 78f37f238b..d93413cbc7 100644 --- a/src/Driving_Accident_Video_Recognition/config.py +++ b/src/Driving_Accident_Video_Recognition/config.py @@ -1,29 +1,88 @@ """ -全局配置文件:所有可配置参数集中管理,新手只需修改这里 +配置模块:支持环境变量/`.env`加载 + 精细化事故判断配置 +优化点:健壮性提升+代码简洁性+可读性增强 """ -# YOLOv8模型配置 -YOLO_MODEL_PATH = "yolov8n.pt" # 轻量化模型(自动下载) -CONFIDENCE_THRESHOLD = 0.5 # 目标检测置信度阈值(0-1) +import os +from typing import Any +from dotenv import load_dotenv -# 检测源配置 -DETECTION_SOURCE = 0 # 0=电脑摄像头,可改为视频路径如"test_accident.mp4" +# 加载.env环境文件(优先读取环境变量,无则用代码默认值) +load_dotenv() -# 事故识别配置 -ACCIDENT_CLASSES = [0, 2, 7] # YOLOv8类别:0=人,2=汽车,7=卡车 -MIN_VEHICLE_COUNT = 2 # 至少2辆车判定为事故 -PERSON_VEHICLE_CONTACT = True # 行人和车辆同时出现判定为事故 -# 帧处理优化配置(低配电脑推荐) -RESIZE_WIDTH = 640 -RESIZE_HEIGHT = 480 +def get_env_config(key: str, default: Any, config_type: type) -> Any: + """ + 统一处理环境变量读取+类型转换(减少重复代码,提升健壮性) + :param key: 环境变量键名 + :param default: 类型匹配的默认值 + :param config_type: 目标类型(int/float/bool/str等) + :return: 转换后的配置值(转换失败自动回退到默认值) + """ + env_value = os.getenv(key) + if env_value is None: + return default + try: + if config_type == bool: + # 布尔类型特殊处理:"True"/"true"转True,其余转False + return env_value.strip().lower() == "true" + return config_type(env_value) + except (ValueError, TypeError): + # 类型转换失败时,回退到默认值 + return default -# 依赖包配置 + +# ====================== YOLO模型配置 ====================== +# YOLO预训练模型路径(默认轻量型yolov8n.pt,适合实时检测) +YOLO_MODEL_PATH = get_env_config("YOLO_MODEL_PATH", "yolov8n.pt", str) +# 检测置信度阈值(0-1,值越高检测越严格,默认0.5平衡精度与召回) +CONFIDENCE_THRESHOLD = get_env_config("CONFIDENCE_THRESHOLD", 0.5, float) + + +# ====================== 检测源配置 ====================== +# 检测源:支持摄像头(整数设备号)或视频文件路径(字符串) +DETECTION_SOURCE = os.getenv("DETECTION_SOURCE", "0") +try: + # 尝试转换为整数(对应摄像头设备号,如0=默认摄像头) + DETECTION_SOURCE = int(DETECTION_SOURCE) +except ValueError: + # 转换失败则视为视频文件路径(保持字符串) + pass + + +# ====================== 事故识别核心配置 ====================== +# 事故检测关注的YOLO类别(0=person/行人、2=car/汽车、7=truck/卡车) +ACCIDENT_CLASSES = [0, 2, 7] +# 多车事故判定阈值:至少检测到N辆车辆才判定为多车事故 +MIN_VEHICLE_COUNT = get_env_config("MIN_VEHICLE_COUNT", 2, int) +# 是否开启“人车接触”事故判定(True=开启,检测到行人+车辆即判定) +PERSON_VEHICLE_CONTACT = get_env_config("PERSON_VEHICLE_CONTACT", True, bool) +# 人车接触距离阈值(像素):行人和车辆框中心距离<该值时,判定为接触 +PERSON_VEHICLE_DISTANCE_THRESHOLD = get_env_config( + "PERSON_VEHICLE_DISTANCE_THRESHOLD", 50, int +) + + +# ====================== 帧处理配置(平衡速度与精度) ====================== +# 检测帧缩放宽度(默认640,YOLO推荐输入尺寸,兼顾速度) +RESIZE_WIDTH = get_env_config("RESIZE_WIDTH", 640, int) +# 检测帧缩放高度(默认480,与宽度配合保持合理比例) +RESIZE_HEIGHT = get_env_config("RESIZE_HEIGHT", 480, int) + + +# ====================== 依赖包配置(自动安装时使用) ====================== REQUIRED_PACKAGES = [ - "ultralytics>=8.0.0", - "opencv-python>=4.8.0", - "numpy>=1.24.0", - "torch>=2.0.0" + "ultralytics>=8.0.0", # YOLOv8核心依赖 + "opencv-python>=4.8.0", # 视频/图像读取、绘制标注 + "numpy>=1.24.0", # 数值计算(坐标/距离运算) + "torch>=2.0.0", # YOLO模型推理(PyTorch后端) + "python-dotenv>=1.0.0" # 加载.env环境变量 ] +# PyPI镜像源(加速国内环境的依赖安装) +PYPI_MIRROR = "https://pypi.tuna.tsinghua.edu.cn/simple" + -# 清华镜像源(加速依赖下载) -PYPI_MIRROR = "https://pypi.tuna.tsinghua.edu.cn/simple" \ No newline at end of file +# ====================== 检测结果输出配置 ====================== +# 是否保存检测结果视频(True=保存,False=不保存) +SAVE_RESULT_VIDEO = get_env_config("SAVE_RESULT_VIDEO", False, bool) +# 检测结果视频保存路径(默认输出到项目根目录) +RESULT_VIDEO_PATH = get_env_config("RESULT_VIDEO_PATH", "detection_result.mp4", str) From 65d544208aeb5b0229336f6e3bd48fc1cf4d5101 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Mon, 22 Dec 2025 10:12:17 +0800 Subject: [PATCH 2/7] Refactor accident video detection tool with enhancements Refactor main program for accident video detection with performance improvements, flexible configurations, and enhanced logging. Added support for recognizing people and vehicles. --- .../main.py | 114 ++++++++++++++---- 1 file changed, 92 insertions(+), 22 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/main.py b/src/Driving_Accident_Video_Recognition/main.py index 073a804427..8419e7e63b 100644 --- a/src/Driving_Accident_Video_Recognition/main.py +++ b/src/Driving_Accident_Video_Recognition/main.py @@ -1,41 +1,111 @@ """ -主程序:驾驶事故视频识别工具 +主程序:驾驶事故视频识别工具(优化版) +优化点:性能提速+灵活配置+规范日志+新增人和小车识别提示 """ import sys import os import argparse +import logging # 新增:日志模块(替代print,支持分级输出) +from config import ( + REQUIRED_PACKAGES, PYPI_MIRROR, DETECTION_SOURCE, + CONFIDENCE_THRESHOLD, ACCIDENT_CLASSES # 新增:引入识别类别配置 +) +from utils.dependencies import install_dependencies +from core.detector import AccidentDetector -# 确保当前目录可被搜索 -current_dir = os.path.dirname(os.path.abspath(__file__)) -sys.path.append(current_dir) +# -------------------------- 新增1:日志初始化(替代print,更灵活) -------------------------- +def init_logger(): + logger = logging.getLogger("AccidentDetection") + logger.setLevel(logging.INFO) + # 控制台输出格式:时间+日志级别+内容 + formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s") + console_handler = logging.StreamHandler() + console_handler.setFormatter(formatter) + logger.addHandler(console_handler) + return logger -# 直接导入同目录的文件(彻底避免模块包问题) -from config import REQUIRED_PACKAGES, PYPI_MIRROR, PERSON_VEHICLE_DISTANCE, ACCIDENT_CONTINUOUS_FRAMES -from dependencies import install_dependencies # 直接导入同目录的dependencies.py -from detector import AccidentDetector +# -------------------------- 新增2:优化命令行参数(更灵活的配置) -------------------------- +def parse_args(logger): + parser = argparse.ArgumentParser(description="驾驶事故视频识别工具(支持动态配置)") + # 基础参数:检测源、语言 + parser.add_argument("--source", "-s", default=DETECTION_SOURCE, + help=f"检测源(0=摄像头/视频路径,默认:{DETECTION_SOURCE})") + parser.add_argument("--language", "-l", default="zh", choices=["zh", "en"], + help="标注语言(zh=中文/en=英文,默认:zh)") + # 新增:性能/配置参数(无需改config.py,直接命令行调整) + parser.add_argument("--skip-deps", "-sd", action="store_true", default=False, + help="跳过依赖检查(已安装依赖时用,提速)") + parser.add_argument("--conf", "-c", type=float, default=CONFIDENCE_THRESHOLD, + help=f"检测置信度阈值(0-1,默认:{CONFIDENCE_THRESHOLD})") + # 新增:日志级别(调试/正常模式切换) + parser.add_argument("--log-level", "-ll", default="INFO", choices=["DEBUG", "INFO", "WARNING"], + help="日志级别(DEBUG=调试/INFO=正常/WARNING=仅警告,默认:INFO)") + + args = parser.parse_args() + # 校验参数合法性(新增:避免无效输入) + if not (0 < args.conf <= 1): + logger.warning(f"置信度{args.conf}无效,自动使用默认值{CONFIDENCE_THRESHOLD}") + args.conf = CONFIDENCE_THRESHOLD + return args +# -------------------------- 优化3:主函数逻辑(减少重复计算+提升健壮性+新增人和小车识别) -------------------------- +def main(): + # 初始化日志 + logger = init_logger() + # 解析参数(并应用日志级别) + args = parse_args(logger) + logger.setLevel(args.log_level) # 动态调整日志级别 -def parse_args(): - parser = argparse.ArgumentParser(description="驾驶事故视频识别") - parser.add_argument("--source", "-s", default=0, help="检测源:0=摄像头/视频路径") - parser.add_argument("--language", "-l", default="zh", choices=["zh", "en"], help="标注语言") - return parser.parse_args() + # -------------------------- 优化4:缓存环境变量操作(减少属性查找,提速) -------------------------- + env = os.environ # 局部变量缓存os.environ,避免循环中重复查找(参考摘要5“缓存属性”) + # 覆盖检测源(命令行优先) + if str(args.source) != str(DETECTION_SOURCE): + # 严谨处理检测源类型:尝试转整数(摄像头),失败则为字符串(视频路径) + try: + env["DETECTION_SOURCE"] = str(int(args.source)) # 摄像头(数字) + except (ValueError, TypeError): + env["DETECTION_SOURCE"] = str(args.source) # 视频路径(字符串) + logger.info(f"检测源已覆盖为:{env['DETECTION_SOURCE']}") + # 覆盖置信度阈值(命令行优先) + if args.conf != CONFIDENCE_THRESHOLD: + env["CONFIDENCE_THRESHOLD"] = str(args.conf) + logger.info(f"置信度阈值已覆盖为:{args.conf}") -def main(): - args = parse_args() try: - print("🚀 启动驾驶事故检测...") - # 安装依赖 - install_dependencies(REQUIRED_PACKAGES, PYPI_MIRROR) - # 启动检测 + logger.info("🚀 启动驾驶事故视频识别工具...") + # -------------------------- 优化5:跳过依赖检查(避免重复安装,提速) -------------------------- + if not args.skip_deps: + install_dependencies(REQUIRED_PACKAGES, PYPI_MIRROR) + else: + logger.info("⚠️ 已跳过依赖检查(--skip-deps生效)") + + # -------------------------- 优化6:简化检测器初始化(减少冗余代码) -------------------------- + logger.info("🔄 初始化事故检测器...") detector = AccidentDetector() + # 新增:提示当前模型支持识别人和小车 + target_classes = {0: "人", 2: "小车"} + supported_targets = [f"{name}(类别ID: {cid})" for cid, name in target_classes.items() if cid in ACCIDENT_CLASSES] + logger.info(f"✅ 检测器初始化完成,当前模型支持识别:{', '.join(supported_targets)}") + logger.info("✅ 开始检测(按Q/ESC退出,画面中会标注识别到的人和小车)") + + # 启动检测(传递语言参数) detector.run_detection(language=args.language) + except KeyboardInterrupt: - print("\n🛑 程序中断") + logger.info("\n🛑 用户强制中断程序") + except Exception as e: + # 新增:DEBUG级别输出详细异常栈,INFO级别只显示错误信息(方便调试) + logger.error(f"\n❌ 程序运行出错:{str(e)}") + if args.log_level == "DEBUG": + import traceback + traceback.print_exc() finally: - print("👋 程序退出") - + logger.info("👋 程序正常退出") if __name__ == "__main__": + # 新增:确保code目录在搜索路径(兼容不同运行方式) + current_dir = os.path.dirname(os.path.abspath(__file__)) + if current_dir not in sys.path: + sys.path.append(current_dir) main() From 6f88d57b0c4aa49e722d6bbbd818d56887cb8f91 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Mon, 22 Dec 2025 10:33:47 +0800 Subject: [PATCH 3/7] Refactor process.py for improved functionality Refactor processing functions for accident video recognition, adding distance calculation and box center retrieval. --- .../process.py | 84 +++++++++++-------- 1 file changed, 48 insertions(+), 36 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/process.py b/src/Driving_Accident_Video_Recognition/process.py index 87df8b1bbd..cba31e2bc1 100644 --- a/src/Driving_Accident_Video_Recognition/process.py +++ b/src/Driving_Accident_Video_Recognition/process.py @@ -1,46 +1,58 @@ """ -辅助处理工具:负责坐标转换、帧处理等通用逻辑 +process.py:辅助函数(处理坐标、距离、标注) """ -import numpy as np -import torch import cv2 +import numpy as np + def process_box_coords(box, scale_x, scale_y): + """将YOLO输出的坐标缩放回原始帧尺寸""" + x1, y1, x2, y2 = box.xyxy[0].tolist() + return ( + int(x1 * scale_x), + int(y1 * scale_y), + int(x2 * scale_x), + int(y2 * scale_y) + ) + + +def get_box_center(x1, y1, x2, y2): + """计算目标框的中心坐标""" + return (int((x1 + x2) / 2), int((y1 + y2) / 2)) + + +def calculate_euclidean_distance(pt1, pt2): + """计算两个点的欧式距离""" + return np.linalg.norm(np.array(pt1) - np.array(pt2)) + + +def draw_annotations(frame, detected_objects, is_accident, language="zh"): """ - 安全处理YOLOv8检测框坐标,解决张量类型错误 - :param box: YOLOv8的检测框对象 - :param scale_x: 宽度缩放比例 - :param scale_y: 高度缩放比例 - :return: 转换后的坐标(x1, y1, x2, y2) - """ - # 兼容张量和numpy数组 - if isinstance(box.xyxy[0], torch.Tensor): - box_xyxy = box.xyxy[0].cpu().numpy() - else: - box_xyxy = np.array(box.xyxy[0]) - # 缩放坐标 - scaled_box = box_xyxy * [scale_x, scale_y, scale_x, scale_y] - # 转换为整数 - return map(int, scaled_box) - -def draw_annotations(frame, detected_objects, is_accident): - """ - 在帧上绘制检测标注和事故警告 - :param frame: 原始视频帧 - :param detected_objects: 检测到的目标列表 - :param is_accident: 是否检测到事故 - :return: 标注后的帧 + 绘制标注(无需扩展库,用拼音+中文注释避免乱码) """ - # 绘制目标检测框 - for (cls_name, x1, y1, x2, y2) in detected_objects: + # 类别映射:拼音+中文(OpenCV默认支持英文/拼音) + class_map = { + "person": "Ren(人)" if language == "zh" else "Person", + "car": "Xiao Che(小车)" if language == "zh" else "Car", + "truck": "Ka Che(卡车)" if language == "zh" else "Truck" + } + + # OpenCV默认英文字体(无需扩展库) + font = cv2.FONT_HERSHEY_SIMPLEX + + # 绘制目标框+标签 + for obj in detected_objects: + cls_name, x1, y1, x2, y2 = obj + display_name = class_map.get(cls_name, cls_name) + # 绘制绿色框 cv2.rectangle(frame, (x1, y1), (x2, y2), (0, 255, 0), 2) - cv2.putText(frame, cls_name, (x1, y1 - 10), - cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 2) - # 绘制事故警告 + # 绘制标签(避免超出画面) + label_y = y1 - 10 if y1 > 20 else y1 + 20 + cv2.putText(frame, display_name, (x1, label_y), font, 0.8, (0, 255, 0), 2) + + # 绘制事故提示(红色) if is_accident: - cv2.putText(frame, "⚠️ 检测到事故!", (50, 50), - cv2.FONT_HERSHEY_SIMPLEX, 1.2, (0, 0, 255), 3) - return frame + accident_text = "Shi Gu!(事故!)" if language == "zh" else "Accident Detected!" + cv2.putText(frame, accident_text, (50, 50), font, 1.2, (0, 0, 255), 3) -# 供外部导入的函数 -__all__ = ["process_box_coords", "draw_annotations"] \ No newline at end of file + return frame From 8ec4153259005c12605508fe3b6a9caa8974d4b0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Mon, 22 Dec 2025 10:34:59 +0800 Subject: [PATCH 4/7] Refactor AccidentDetector for improved functionality Enhance accident detection with video saving and FPS display. --- .../detector.py | 193 ++++++++++++------ 1 file changed, 132 insertions(+), 61 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/detector.py b/src/Driving_Accident_Video_Recognition/detector.py index 4155840ab8..d281ccaf87 100644 --- a/src/Driving_Accident_Video_Recognition/detector.py +++ b/src/Driving_Accident_Video_Recognition/detector.py @@ -1,121 +1,192 @@ """ -事故检测器核心类:负责模型加载、检测流程执行 +检测器模块:精准事故判断+视频保存+帧率显示(无报错版) """ import sys import cv2 +import time from ultralytics import YOLO from config import ( YOLO_MODEL_PATH, CONFIDENCE_THRESHOLD, ACCIDENT_CLASSES, - MIN_VEHICLE_COUNT, PERSON_VEHICLE_CONTACT, - RESIZE_WIDTH, RESIZE_HEIGHT, DETECTION_SOURCE + MIN_VEHICLE_COUNT, PERSON_VEHICLE_CONTACT, PERSON_VEHICLE_DISTANCE_THRESHOLD, + RESIZE_WIDTH, RESIZE_HEIGHT, DETECTION_SOURCE, + SAVE_RESULT_VIDEO, RESULT_VIDEO_PATH ) -from core.process import process_box_coords, draw_annotations +from core.process import ( + process_box_coords, get_box_center, calculate_euclidean_distance, draw_annotations +) + class AccidentDetector: def __init__(self): - """初始化检测器,加载YOLOv8模型""" - self.model = None - self.accident_detected = False - self._load_model() + self.model = None # YOLO模型对象 + self.accident_detected = False # 是否检测到事故 + self.video_writer = None # 视频写入器(保存检测结果) + # 帧率计算(滑动平均,避免波动) + self.fps_history = [] + self.prev_time = time.time() + + self._load_model() # 初始化时加载模型 def _load_model(self): - """私有方法:加载模型,包含重试逻辑""" + """加载YOLO模型(增加兜底逻辑)""" + print("🔄 加载YOLOv8检测模型...") try: - print("🔄 正在加载YOLOv8模型(首次运行会自动下载)...") self.model = YOLO(YOLO_MODEL_PATH) - print("✅ YOLOv8模型加载成功") + print(f"✅ 模型加载成功:{YOLO_MODEL_PATH}") except Exception as e: - print(f"❌ 模型加载失败:{e}") - # 重试加载模型 + print(f"⚠️ 指定模型加载失败,尝试默认轻量模型yolov8n.pt...") try: - print("🔄 尝试重新下载模型...") self.model = YOLO("yolov8n.pt") - print("✅ 模型重新加载成功") + print("✅ 兜底模型(yolov8n.pt)加载成功") except Exception as e2: - print(f"❌ 模型重新加载失败:{e2}") + print(f"❌ 模型加载失败:{e2},程序退出") sys.exit(1) - def detect_frame(self, frame): - """处理单帧,返回标注后的帧和是否检测到事故""" + def _init_video_writer(self, frame): + """初始化视频写入器(增加路径检查)""" + if not SAVE_RESULT_VIDEO: + return + height, width = frame.shape[:2] + fourcc = cv2.VideoWriter_fourcc(*"mp4v") + # 自动创建保存目录(避免路径不存在) + save_dir = "/".join(RESULT_VIDEO_PATH.split("/")[:-1]) + if save_dir and not cv2.os.path.exists(save_dir): + cv2.os.makedirs(save_dir) + # 初始化写入器 + self.video_writer = cv2.VideoWriter(RESULT_VIDEO_PATH, fourcc, 30.0, (width, height)) + if not self.video_writer.isOpened(): + print(f"⚠️ 无法保存视频到{RESULT_VIDEO_PATH},跳过保存") + self.video_writer = None + + def _calculate_accident(self, detected_objects): + """精准判断事故:多车/人车接触""" + persons = [obj for obj in detected_objects if obj[0] == "person"] + vehicles = [obj for obj in detected_objects if obj[0] in ["car", "truck"]] + + # 条件1:车辆数量≥配置阈值 + if len(vehicles) >= MIN_VEHICLE_COUNT: + return True + # 条件2:行人和车辆距离≤阈值 + if PERSON_VEHICLE_CONTACT and len(persons) >= 1 and len(vehicles) >= 1: + p_centers = [get_box_center(*obj[1:]) for obj in persons] + v_centers = [get_box_center(*obj[1:]) for obj in vehicles] + for p in p_centers: + for v in v_centers: + if calculate_euclidean_distance(p, v) <= PERSON_VEHICLE_DISTANCE_THRESHOLD: + return True + return False + + def detect_frame(self, frame, language="zh"): + """处理单帧:检测+标注+帧率计算""" detected_objects = [] + current_frame = frame.copy() + try: - # 缩放帧提升速度 - frame_resized = cv2.resize(frame, (RESIZE_WIDTH, RESIZE_HEIGHT)) - # YOLOv8推理 - results = self.model(frame_resized, conf=CONFIDENCE_THRESHOLD) + # 缩放帧(适配YOLO输入) + frame_resized = cv2.resize(current_frame, (RESIZE_WIDTH, RESIZE_HEIGHT)) + # 模型推理(关闭冗余日志) + results = self.model(frame_resized, conf=CONFIDENCE_THRESHOLD, verbose=False) # 解析检测结果 for r in results: - if hasattr(r, 'boxes') and r.boxes is not None: - for box in r.boxes: - if not hasattr(box, 'cls') or box.cls is None: - continue - cls_idx = int(box.cls[0]) - if cls_idx in ACCIDENT_CLASSES: - cls_name = self.model.names[cls_idx] - # 处理坐标 - scale_x = frame.shape[1] / RESIZE_WIDTH - scale_y = frame.shape[0] / RESIZE_HEIGHT - x1, y1, x2, y2 = process_box_coords(box, scale_x, scale_y) - detected_objects.append((cls_name, x1, y1, x2, y2)) + if not hasattr(r, "boxes") or r.boxes is None: + continue + for box in r.boxes: + if not hasattr(box, "cls") or box.cls is None: + continue + cls_idx = int(box.cls[0]) + if cls_idx in ACCIDENT_CLASSES: + cls_name = self.model.names[cls_idx] + # 坐标缩放回原始帧 + scale_x = current_frame.shape[1] / RESIZE_WIDTH + scale_y = current_frame.shape[0] / RESIZE_HEIGHT + x1, y1, x2, y2 = process_box_coords(box, scale_x, scale_y) + detected_objects.append((cls_name, x1, y1, x2, y2)) # 判断事故 - person_count = sum(1 for obj in detected_objects if obj[0] == "person") - vehicle_count = sum(1 for obj in detected_objects if obj[0] in ["car", "truck"]) - is_accident = (vehicle_count >= MIN_VEHICLE_COUNT) or (person_count >= 1 and vehicle_count >= 1 and PERSON_VEHICLE_CONTACT) - self.accident_detected = is_accident - + self.accident_detected = self._calculate_accident(detected_objects) # 绘制标注 - frame = draw_annotations(frame, detected_objects, is_accident) + current_frame = draw_annotations(current_frame, detected_objects, self.accident_detected, language) + + # 计算滑动平均帧率 + current_time = time.time() + self.fps_history.append(1 / (current_time - self.prev_time)) + self.prev_time = current_time + # 只保留最近10帧的帧率(避免波动) + if len(self.fps_history) > 10: + self.fps_history.pop(0) + avg_fps = int(sum(self.fps_history) / len(self.fps_history)) if self.fps_history else 0 + # 绘制帧率 + cv2.putText(current_frame, f"FPS: {avg_fps}", (50, 100), + cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 0, 0), 2) + + # 保存视频帧 + if self.video_writer: + self.video_writer.write(current_frame) except Exception as e: - print(f"⚠️ 帧处理出现小错误:{e},继续运行...") + print(f"⚠️ 帧处理错误:{e},继续运行...") - return frame, self.accident_detected + return current_frame, self.accident_detected - def run_detection(self): - """启动检测流程,包含完善的容错逻辑""" - # 多次尝试打开检测源 + def run_detection(self, language="zh"): + """启动检测流程:打开摄像头/视频+逐帧处理""" + # 打开检测源(重试3次) cap = None - for i in range(3): + for retry in range(3): cap = cv2.VideoCapture(DETECTION_SOURCE) if cap.isOpened(): + print(f"✅ 第{retry+1}次打开检测源成功") break - print(f"⚠️ 第{i+1}次打开检测源失败,重试中...") - cv2.waitKey(1000) + print(f"⚠️ 第{retry+1}次打开检测源失败,1秒后重试...") + time.sleep(1) + # 兜底:打开默认摄像头 if not cap or not cap.isOpened(): - print(f"❌ 无法打开检测源:{DETECTION_SOURCE}") - # 强制切换为摄像头 - print("🔄 强制切换为电脑摄像头...") + print(f"❌ 目标检测源{DETECTION_SOURCE}无法打开,尝试默认摄像头(0)...") cap = cv2.VideoCapture(0) if not cap.isOpened(): - print("❌ 摄像头也无法打开,请检查设备") + print("❌ 所有检测源均无法打开,程序退出") sys.exit(1) - print("✅ 检测源打开成功,开始实时检测(按Q/ESC键退出)") - print("💡 提示:检测到2辆车或行人和车辆同时出现时,显示红色警告") + print("✅ 检测源打开成功(按Q/ESC退出)") + print(f"💡 配置:行人车辆距离阈值{PERSON_VEHICLE_DISTANCE_THRESHOLD}像素") + + # 初始化视频写入器(读取第一帧) + ret, first_frame = cap.read() + if ret: + self._init_video_writer(first_frame) # 逐帧处理 while True: ret, frame = cap.read() if not ret: - print("🔚 视频/摄像头流结束") + print("🔚 视频流读取完毕,结束检测") break - frame, _ = self.detect_frame(frame) - cv2.imshow("驾驶事故检测(按Q退出)", frame) + # 处理单帧 + processed_frame, _ = self.detect_frame(frame, language) + cv2.imshow("驾驶事故检测", processed_frame) # 退出逻辑 key = cv2.waitKey(1) & 0xFF - if key == ord('q') or key == 27: + if key == ord("q") or key == 27: print("🛑 用户手动退出") break # 释放资源 cap.release() + if self.video_writer: + self.video_writer.release() + print(f"✅ 检测结果已保存到{RESULT_VIDEO_PATH}") cv2.destroyAllWindows() - print(f"\n📊 检测总结:是否检测到事故 → {'✅ 是' if self.accident_detected else '❌ 否'}") -# 供外部导入的类 -__all__ = ["AccidentDetector"] \ No newline at end of file + # 检测总结 + avg_fps = int(sum(self.fps_history) / len(self.fps_history)) if self.fps_history else 0 + print(f"\n📊 检测总结:") + print(f" - 是否检测到事故 → {'✅ 是' if self.accident_detected else '❌ 否'}") + print(f" - 平均处理帧率 → {avg_fps} FPS") + + +# 供外部导入 +__all__ = ["AccidentDetector"] From 3ab85a3236e9f00c90cfccc2cb1f50a96aa5c566 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Mon, 22 Dec 2025 18:00:44 +0800 Subject: [PATCH 5/7] Refactor driving accident video recognition tool Refactor main program for driving accident video recognition tool with enhanced performance, flexible configuration, and improved logging. Added support for detecting people and vehicles. From 1e143e896f265a572548ca16d5b4571ff142b5c6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=BB=84=E5=87=AF?= <2684756428@qq.com> Date: Mon, 22 Dec 2025 20:18:05 +0800 Subject: [PATCH 6/7] Refactor accident video recognition tool for performance MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 优化了驾驶事故视频识别工具的性能和灵活性,支持命令行动态配置,替换print为日志模块,增强了用户交互体验和程序健壮性。 --- .../main.py | 210 +++++++++++++----- 1 file changed, 157 insertions(+), 53 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/main.py b/src/Driving_Accident_Video_Recognition/main.py index 8419e7e63b..af6586c286 100644 --- a/src/Driving_Accident_Video_Recognition/main.py +++ b/src/Driving_Accident_Video_Recognition/main.py @@ -1,111 +1,215 @@ """ 主程序:驾驶事故视频识别工具(优化版) -优化点:性能提速+灵活配置+规范日志+新增人和小车识别提示 +优化点说明: +1. 性能提速:跳过重复依赖检查、缓存环境变量减少属性查找、简化检测器初始化逻辑 +2. 灵活配置:支持命令行动态调整检测源/置信度/日志级别/语言,无需修改配置文件 +3. 规范日志:替换print为logging模块,支持分级输出(DEBUG/INFO/WARNING),便于调试和生产环境使用 +4. 交互优化:新增人和小车识别提示,明确告知用户当前模型支持的识别类别 +5. 健壮性提升:参数合法性校验、异常捕获并分级输出、兼容不同运行路径 """ -import sys -import os -import argparse -import logging # 新增:日志模块(替代print,支持分级输出) +# 系统内置模块:基础功能支撑 +import sys # 系统路径、退出等核心操作 +import os # 环境变量、文件路径等操作系统交互 +import argparse # 命令行参数解析工具 +import logging # 日志模块(替代print,支持分级输出、格式化、持久化等) + +# 自定义模块/配置:项目核心配置和工具 from config import ( - REQUIRED_PACKAGES, PYPI_MIRROR, DETECTION_SOURCE, - CONFIDENCE_THRESHOLD, ACCIDENT_CLASSES # 新增:引入识别类别配置 + REQUIRED_PACKAGES, # 项目必需的依赖包列表(如ultralytics/opencv-python等) + PYPI_MIRROR, # PyPI镜像源(国内默认清华镜像,提速依赖安装) + DETECTION_SOURCE, # 默认检测源(0=本地摄像头,也可传视频文件路径) + CONFIDENCE_THRESHOLD, # 默认检测置信度阈值(过滤低置信度的识别结果) + ACCIDENT_CLASSES # 事故识别核心类别(0=人,2=小车,7=卡车等) ) -from utils.dependencies import install_dependencies -from core.detector import AccidentDetector +from utils.dependencies import install_dependencies # 依赖自动安装工具函数 +from core.detector import AccidentDetector # 事故检测器核心类(封装YOLO模型、检测逻辑) -# -------------------------- 新增1:日志初始化(替代print,更灵活) -------------------------- +# -------------------------- 核心工具函数1:日志初始化(替代print,更专业、灵活) -------------------------- def init_logger(): + """ + 初始化日志系统(核心作用:统一日志格式、支持分级输出) + 返回值: + logger: 配置好的日志实例,可调用logger.info/debug/warning/error输出不同级别日志 + 日志格式:时间戳 - 日志级别 - 日志内容(例如:2025-12-22 10:00:00 - INFO - 启动程序) + """ + # 创建日志器实例,命名为"AccidentDetection"(便于多模块日志区分) logger = logging.getLogger("AccidentDetection") + # 设置默认日志级别为INFO(低于INFO的日志不会输出,如DEBUG) logger.setLevel(logging.INFO) - # 控制台输出格式:时间+日志级别+内容 - formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s") + + # 避免重复添加处理器(多次调用该函数时防止日志重复输出) + if logger.handlers: + return logger + + # 定义日志输出格式:时间+级别+内容 + formatter = logging.Formatter( + "%(asctime)s - %(levelname)s - %(message)s", # 格式字符串 + datefmt="%Y-%m-%d %H:%M:%S" # 时间格式(可读性更强) + ) + + # 创建控制台处理器(日志输出到终端) console_handler = logging.StreamHandler() - console_handler.setFormatter(formatter) + console_handler.setFormatter(formatter) # 绑定格式 + + # 将处理器添加到日志器 logger.addHandler(console_handler) + return logger -# -------------------------- 新增2:优化命令行参数(更灵活的配置) -------------------------- +# -------------------------- 核心工具函数2:命令行参数解析(灵活配置,无需改代码) -------------------------- def parse_args(logger): + """ + 解析命令行参数(核心作用:让用户通过命令行动态配置参数,提升工具灵活性) + 参数: + logger: 日志实例(用于输出参数校验警告) + 返回值: + args: 解析后的参数对象,可通过args.xxx访问具体参数 + """ + # 创建参数解析器,添加工具描述(--help时显示) parser = argparse.ArgumentParser(description="驾驶事故视频识别工具(支持动态配置)") - # 基础参数:检测源、语言 - parser.add_argument("--source", "-s", default=DETECTION_SOURCE, - help=f"检测源(0=摄像头/视频路径,默认:{DETECTION_SOURCE})") - parser.add_argument("--language", "-l", default="zh", choices=["zh", "en"], - help="标注语言(zh=中文/en=英文,默认:zh)") - # 新增:性能/配置参数(无需改config.py,直接命令行调整) - parser.add_argument("--skip-deps", "-sd", action="store_true", default=False, - help="跳过依赖检查(已安装依赖时用,提速)") - parser.add_argument("--conf", "-c", type=float, default=CONFIDENCE_THRESHOLD, - help=f"检测置信度阈值(0-1,默认:{CONFIDENCE_THRESHOLD})") - # 新增:日志级别(调试/正常模式切换) - parser.add_argument("--log-level", "-ll", default="INFO", choices=["DEBUG", "INFO", "WARNING"], - help="日志级别(DEBUG=调试/INFO=正常/WARNING=仅警告,默认:INFO)") + # 1. 基础配置参数:检测源(摄像头/视频文件) + parser.add_argument( + "--source", "-s", # 参数名(长/短格式) + default=DETECTION_SOURCE, # 默认值(从config.py读取) + help=f"检测源(0=本地摄像头/视频文件绝对路径,默认值:{DETECTION_SOURCE})" + ) + + # 2. 界面配置参数:标注语言(中文/英文) + parser.add_argument( + "--language", "-l", + default="zh", # 默认中文 + choices=["zh", "en"], # 限制可选值(避免无效输入) + help="标注语言(zh=中文/en=英文,默认:zh)" + ) + + # 3. 性能优化参数:跳过依赖检查(已安装依赖时提速) + parser.add_argument( + "--skip-deps", "-sd", + action="store_true", # 无需传值,加该参数则为True + default=False, + help="跳过依赖检查(已确认安装所有依赖时使用,可大幅提升启动速度)" + ) + + # 4. 检测精度参数:置信度阈值(过滤低置信度结果) + parser.add_argument( + "--conf", "-c", + type=float, # 参数类型(浮点型) + default=CONFIDENCE_THRESHOLD, + help=f"检测置信度阈值(范围0-1,值越高越严格,默认:{CONFIDENCE_THRESHOLD})" + ) + + # 5. 调试配置参数:日志级别(控制输出详细程度) + parser.add_argument( + "--log-level", "-ll", + default="INFO", # 默认只输出INFO及以上级别 + choices=["DEBUG", "INFO", "WARNING"], # 可选级别 + help="日志级别(DEBUG=调试模式/INFO=正常模式/WARNING=仅警告,默认:INFO)" + ) + + # 解析命令行传入的参数 args = parser.parse_args() - # 校验参数合法性(新增:避免无效输入) + + # 关键校验:置信度阈值合法性(必须在0-1之间) if not (0 < args.conf <= 1): - logger.warning(f"置信度{args.conf}无效,自动使用默认值{CONFIDENCE_THRESHOLD}") + # 输出警告日志,自动回退到默认值 + logger.warning(f"输入的置信度{args.conf}无效(需0 Date: Tue, 23 Dec 2025 17:00:19 +0800 Subject: [PATCH 7/7] Optimize detector with accident type and confidence Enhance accident detection with type differentiation, confidence scoring, and target counting. --- .../detector.py | 98 +++++++++++-------- 1 file changed, 59 insertions(+), 39 deletions(-) diff --git a/src/Driving_Accident_Video_Recognition/detector.py b/src/Driving_Accident_Video_Recognition/detector.py index d281ccaf87..f840b486de 100644 --- a/src/Driving_Accident_Video_Recognition/detector.py +++ b/src/Driving_Accident_Video_Recognition/detector.py @@ -1,5 +1,5 @@ """ -检测器模块:精准事故判断+视频保存+帧率显示(无报错版) +检测器模块:精准事故判断+视频保存+帧率显示(优化版:新增事故类型区分+置信度+目标计数) """ import sys import cv2 @@ -15,7 +15,6 @@ process_box_coords, get_box_center, calculate_euclidean_distance, draw_annotations ) - class AccidentDetector: def __init__(self): self.model = None # YOLO模型对象 @@ -24,7 +23,6 @@ def __init__(self): # 帧率计算(滑动平均,避免波动) self.fps_history = [] self.prev_time = time.time() - self._load_model() # 初始化时加载模型 def _load_model(self): @@ -59,35 +57,38 @@ def _init_video_writer(self, frame): self.video_writer = None def _calculate_accident(self, detected_objects): - """精准判断事故:多车/人车接触""" + """精准判断事故类型:返回None/多车事故/人车接触事故""" persons = [obj for obj in detected_objects if obj[0] == "person"] vehicles = [obj for obj in detected_objects if obj[0] in ["car", "truck"]] - - # 条件1:车辆数量≥配置阈值 + + # 条件1:多车事故(车辆数量≥配置阈值) if len(vehicles) >= MIN_VEHICLE_COUNT: - return True - # 条件2:行人和车辆距离≤阈值 + return "multi_vehicle" + # 条件2:人车接触事故(行人和车辆距离≤阈值) if PERSON_VEHICLE_CONTACT and len(persons) >= 1 and len(vehicles) >= 1: p_centers = [get_box_center(*obj[1:]) for obj in persons] v_centers = [get_box_center(*obj[1:]) for obj in vehicles] for p in p_centers: for v in v_centers: if calculate_euclidean_distance(p, v) <= PERSON_VEHICLE_DISTANCE_THRESHOLD: - return True - return False + return "person_vehicle" + # 无事故 + return None def detect_frame(self, frame, language="zh"): - """处理单帧:检测+标注+帧率计算""" + """处理单帧:新增目标计数+置信度显示+事故类型区分""" detected_objects = [] current_frame = frame.copy() - + # 新增:目标数量统计(人、小车、卡车) + target_count = {"person": 0, "car": 0, "truck": 0} + try: # 缩放帧(适配YOLO输入) frame_resized = cv2.resize(current_frame, (RESIZE_WIDTH, RESIZE_HEIGHT)) # 模型推理(关闭冗余日志) results = self.model(frame_resized, conf=CONFIDENCE_THRESHOLD, verbose=False) - - # 解析检测结果 + + # 解析检测结果(新增置信度提取) for r in results: if not hasattr(r, "boxes") or r.boxes is None: continue @@ -97,36 +98,64 @@ def detect_frame(self, frame, language="zh"): cls_idx = int(box.cls[0]) if cls_idx in ACCIDENT_CLASSES: cls_name = self.model.names[cls_idx] + # 新增:获取检测置信度(保留2位小数) + conf = round(float(box.conf[0]), 2) # 坐标缩放回原始帧 scale_x = current_frame.shape[1] / RESIZE_WIDTH scale_y = current_frame.shape[0] / RESIZE_HEIGHT x1, y1, x2, y2 = process_box_coords(box, scale_x, scale_y) - detected_objects.append((cls_name, x1, y1, x2, y2)) - - # 判断事故 - self.accident_detected = self._calculate_accident(detected_objects) - # 绘制标注 - current_frame = draw_annotations(current_frame, detected_objects, self.accident_detected, language) - - # 计算滑动平均帧率 + detected_objects.append((cls_name, conf, x1, y1, x2, y2)) # 新增conf参数 + # 统计目标数量 + target_count[cls_name] += 1 + + # 判定事故类型(替代原布尔值判断) + accident_type = self._calculate_accident(detected_objects) + self.accident_detected = accident_type is not None + + # 绘制标注(适配新增的置信度和事故类型) + font = cv2.FONT_HERSHEY_SIMPLEX + # 1. 绘制目标框+标签(含置信度) + for obj in detected_objects: + cls_name, conf, x1, y1, x2, y2 = obj + # 类别名称映射(保留原逻辑) + class_map = { + "person": "Ren(人)" if language == "zh" else "Person", + "car": "Xiao Che(小车)" if language == "zh" else "Car", + "truck": "Ka Che(卡车)" if language == "zh" else "Truck" + } + display_name = f"{class_map.get(cls_name, cls_name)}({conf})" # 新增置信度显示 + # 绘制绿色框(原逻辑不变) + cv2.rectangle(current_frame, (x1, y1), (x2, y2), (0, 255, 0), 2) + # 绘制标签(避免超出画面) + label_y = y1 - 10 if y1 > 20 else y1 + 20 + cv2.putText(current_frame, display_name, (x1, label_y), font, 0.8, (0, 255, 0), 2) + + # 2. 绘制事故提示(按类型区分颜色) + if accident_type == "multi_vehicle": + accident_text = "Duo Che Shi Gu!(多车事故!)" if language == "zh" else "Multi-Vehicle Accident!" + cv2.putText(current_frame, accident_text, (50, 50), font, 1.2, (0, 255, 255), 3) # 黄色 + elif accident_type == "person_vehicle": + accident_text = "Ren Che Jie Chu!(人车接触!)" if language == "zh" else "Person-Vehicle Contact!" + cv2.putText(current_frame, accident_text, (50, 50), font, 1.2, (0, 0, 255), 3) # 红色 + + # 3. 绘制目标数量统计(新增) + count_text = f"Ren: {target_count['person']} | Xiao Che: {target_count['car']} | Ka Che: {target_count['truck']}" if language == "zh" else f"Person: {target_count['person']} | Car: {target_count['car']} | Truck: {target_count['truck']}" + cv2.putText(current_frame, count_text, (50, 150), font, 0.8, (255, 255, 0), 2) # 青色 + + # 4. 绘制帧率(调整位置避免重叠) current_time = time.time() self.fps_history.append(1 / (current_time - self.prev_time)) self.prev_time = current_time - # 只保留最近10帧的帧率(避免波动) if len(self.fps_history) > 10: self.fps_history.pop(0) avg_fps = int(sum(self.fps_history) / len(self.fps_history)) if self.fps_history else 0 - # 绘制帧率 - cv2.putText(current_frame, f"FPS: {avg_fps}", (50, 100), - cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 0, 0), 2) - - # 保存视频帧 + cv2.putText(current_frame, f"FPS: {avg_fps}", (50, 100), font, 1, (255, 0, 0), 2) + + # 保存视频帧(原逻辑不变) if self.video_writer: self.video_writer.write(current_frame) - except Exception as e: print(f"⚠️ 帧处理错误:{e},继续运行...") - return current_frame, self.accident_detected def run_detection(self, language="zh"): @@ -140,7 +169,6 @@ def run_detection(self, language="zh"): break print(f"⚠️ 第{retry+1}次打开检测源失败,1秒后重试...") time.sleep(1) - # 兜底:打开默认摄像头 if not cap or not cap.isOpened(): print(f"❌ 目标检测源{DETECTION_SOURCE}无法打开,尝试默认摄像头(0)...") @@ -148,45 +176,37 @@ def run_detection(self, language="zh"): if not cap.isOpened(): print("❌ 所有检测源均无法打开,程序退出") sys.exit(1) - print("✅ 检测源打开成功(按Q/ESC退出)") print(f"💡 配置:行人车辆距离阈值{PERSON_VEHICLE_DISTANCE_THRESHOLD}像素") - # 初始化视频写入器(读取第一帧) ret, first_frame = cap.read() if ret: self._init_video_writer(first_frame) - # 逐帧处理 while True: ret, frame = cap.read() if not ret: print("🔚 视频流读取完毕,结束检测") break - # 处理单帧 processed_frame, _ = self.detect_frame(frame, language) cv2.imshow("驾驶事故检测", processed_frame) - # 退出逻辑 key = cv2.waitKey(1) & 0xFF if key == ord("q") or key == 27: print("🛑 用户手动退出") break - # 释放资源 cap.release() if self.video_writer: self.video_writer.release() print(f"✅ 检测结果已保存到{RESULT_VIDEO_PATH}") cv2.destroyAllWindows() - # 检测总结 avg_fps = int(sum(self.fps_history) / len(self.fps_history)) if self.fps_history else 0 print(f"\n📊 检测总结:") print(f" - 是否检测到事故 → {'✅ 是' if self.accident_detected else '❌ 否'}") print(f" - 平均处理帧率 → {avg_fps} FPS") - # 供外部导入 __all__ = ["AccidentDetector"]