摘要

随着人工智能技术的快速发展,基于深度学习的计算机视觉技术在安防监控、公共卫生等领域展现出巨大潜力。吸烟检测作为公共安全和健康管理的重要组成部分,传统检测方法往往效率低下且准确率有限。本文将详细介绍基于YOLO系列算法(YOLOv5/v6/v7/v8)的吸烟检测系统的设计与实现,涵盖算法原理、数据集构建、模型训练、系统集成及可视化界面开发等完整流程。通过本系统,可实现实时、高效的吸烟行为检测,为公共场所禁烟管理提供智能化解决方案。

目录

摘要

1. 引言

1.1 研究背景与意义

1.2 YOLO算法发展概述

1.3 系统设计目标

2. 吸烟检测数据集构建

2.1 数据集来源与采集

2.2 数据标注规范

2.3 数据增强策略

3. YOLO算法原理与改进

3.1 YOLOv5架构详解

3.2 YOLOv6与YOLOv7的改进

3.3 YOLOv8的创新特性

4. 系统设计与实现

4.1 系统架构设计

4.2 完整代码实现

4.2.1 项目结构

4.2.2 配置文件

4.2.3 主训练代码

4.2.4 检测推理代码

4.2.5 PySide6图形界面

4.2.6 训练配置文件


1. 引言

1.1 研究背景与意义

吸烟不仅危害个人健康,在公共场所吸烟还可能引发火灾风险,影响他人健康。传统的人工监控方式存在效率低、成本高、易漏检等问题。基于深度学习的吸烟检测系统能够实现7×24小时不间断自动监控,及时发现并预警吸烟行为,对提升公共场所安全管理水平具有重要意义。

1.2 YOLO算法发展概述

YOLO(You Only Look Once)系列算法自2015年问世以来,以其卓越的实时检测性能受到广泛关注。从最初的YOLOv1到最新的YOLOv8,每一代都在精度和速度上取得了显著提升。YOLOv5凭借其简单易用的特性成为工业界最受欢迎的目标检测框架之一,YOLOv6、YOLOv7、YOLOv8则在网络结构、训练策略等方面进行了多项创新。

1.3 系统设计目标

本系统旨在实现以下目标:

  • 实现对吸烟行为的实时准确检测

  • 支持多种YOLO版本模型训练与推理

  • 提供友好的可视化操作界面

  • 具备良好的可扩展性和部署便利性

2. 吸烟检测数据集构建

2.1 数据集来源与采集

吸烟检测数据集可从多个公开数据集和实际场景采集获得:

公开数据集:

  1. Smoking Dataset from Kaggle: 包含多种场景下的吸烟图像

  2. Custom Smoking Dataset: 通过公开数据收集整理的吸烟检测数据集

  3. Real-world Surveillance Data: 实际监控场景采集数据(需合规处理)

数据采集建议:

  • 多场景覆盖:室内、室外、白天、夜晚等不同环境

  • 多角度采集:正面、侧面、背面等不同角度

  • 多样化对象:不同年龄、性别、衣着特征

  • 合规性注意:尊重隐私,仅使用公开或授权数据

2.2 数据标注规范

使用LabelImg或CVAT等工具进行标注,标注类别为:

  • smoking: 吸烟行为(手持香烟且正在吸烟)

  • cigarette: 香烟(未点燃或手持但未吸烟)

  • person: 人物(用于辅助检测)

2.3 数据增强策略

为提高模型泛化能力,采用以下数据增强技术:

python

# 数据增强配置示例
augmentation = {
    'hsv_h': 0.015,  # 色调增强
    'hsv_s': 0.7,    # 饱和度增强
    'hsv_v': 0.4,    # 明度增强
    'rotate': 10,    # 旋转角度
    'translate': 0.1, # 平移
    'scale': 0.5,    # 缩放
    'shear': 0.0,    # 剪切
    'perspective': 0.0, # 透视变换
    'flipud': 0.0,   # 上下翻转
    'fliplr': 0.5,   # 左右翻转
    'mosaic': 1.0,   # 马赛克增强
    'mixup': 0.0,    # MixUp增强
}

3. YOLO算法原理与改进

3.1 YOLOv5架构详解

YOLOv5采用CSPDarknet53作为Backbone,SPP和PANet作为Neck,三个检测头作为Head。其创新点包括:

  • 自适应锚框计算

  • 自适应图片缩放

  • Mosaic数据增强

  • 自动学习数据增强参数

3.2 YOLOv6与YOLOv7的改进

YOLOv6引入RepVGG风格的重参数化Backbone,YOLOv7提出扩展高效层聚合网络和标签分配策略。

3.3 YOLOv8的创新特性

YOLOv8采用新的骨干网络和检测头设计,取消锚框机制,使用无锚检测,简化训练流程。

4. 系统设计与实现

4.1 系统架构设计

系统采用模块化设计,包含以下核心模块:

  1. 数据预处理模块

  2. 模型训练模块

  3. 推理检测模块

  4. 可视化界面模块

  5. 结果管理模块

4.2 完整代码实现

4.2.1 项目结构

text

smoking_detection_system/
│
├── data/                    # 数据集目录
│   ├── images/             # 图像数据
│   │   ├── train/         # 训练集
│   │   └── val/           # 验证集
│   └── labels/             # 标注数据
│
├── models/                 # 模型定义
│   ├── yolov5/            # YOLOv5模型
│   ├── yolov6/            # YOLOv6模型
│   ├── yolov7/            # YOLOv7模型
│   └── yolov8/            # YOLOv8模型
│
├── utils/                  # 工具函数
│   ├── datasets.py        # 数据集处理
│   ├── augmentations.py   # 数据增强
│   ├── metrics.py         # 评估指标
│   └── visualization.py   # 可视化工具
│
├── train.py               # 训练脚本
├── detect.py              # 检测脚本
├── evaluate.py            # 评估脚本
├── requirements.txt       # 依赖包
└── README.md             # 项目说明
4.2.2 配置文件

data.yaml 数据集配置文件:

yaml

# 数据集路径
path: ../datasets/smoking  # 数据集根目录
train: images/train  # 训练集路径
val: images/val      # 验证集路径
test: images/test    # 测试集路径

# 类别数量
nc: 3  # 类别数: smoking, cigarette, person

# 类别名称
names: 
  0: smoking
  1: cigarette
  2: person
4.2.3 主训练代码

python

"""
基于YOLO的吸烟检测系统训练脚本
支持YOLOv5/v6/v7/v8
作者:深度学习专家
日期:2024年1月
"""

import os
import sys
import argparse
import yaml
import torch
import numpy as np
from pathlib import Path
from datetime import datetime

# 添加项目根目录到路径
FILE = Path(__file__).resolve()
ROOT = FILE.parents[0]
if str(ROOT) not in sys.path:
    sys.path.append(str(ROOT))

from models.yolov5.train import train as train_yolov5
from models.yolov6.train import train as train_yolov6
from models.yolov7.train import train as train_yolov7
from models.yolov8.train import train as train_yolov8
from utils.general import colorstr, check_dataset, increment_path

def parse_opt():
    """解析命令行参数"""
    parser = argparse.ArgumentParser(description='YOLO吸烟检测训练脚本')
    parser.add_argument('--model', type=str, default='yolov8', 
                       choices=['yolov5', 'yolov6', 'yolov7', 'yolov8'],
                       help='选择YOLO版本')
    parser.add_argument('--data', type=str, default='data/smoking.yaml', 
                       help='数据集配置文件路径')
    parser.add_argument('--epochs', type=int, default=100, 
                       help='训练轮数')
    parser.add_argument('--batch-size', type=int, default=16, 
                       help='批次大小')
    parser.add_argument('--img-size', nargs='+', type=int, default=[640, 640], 
                       help='输入图像尺寸 [height, width]')
    parser.add_argument('--device', default='', 
                       help='训练设备, cuda device, i.e. 0 or 0,1,2,3 or cpu')
    parser.add_argument('--workers', type=int, default=8, 
                       help='数据加载线程数')
    parser.add_argument('--weights', type=str, default='', 
                       help='初始权重路径')
    parser.add_argument('--project', default='runs/train', 
                       help='保存结果的根目录')
    parser.add_argument('--name', default='exp', 
                       help='实验名称')
    parser.add_argument('--exist-ok', action='store_true', 
                       help='允许覆盖现有实验')
    parser.add_argument('--hyp', type=str, default='', 
                       help='超参数文件路径')
    parser.add_argument('--patience', type=int, default=100, 
                       help='早停耐心值')
    parser.add_argument('--save-period', type=int, default=-1, 
                       help='每N轮保存一次检查点')
    parser.add_argument('--seed', type=int, default=42, 
                       help='随机种子')
    
    return parser.parse_args()

def setup_environment(opt):
    """设置训练环境"""
    # 设置随机种子
    torch.manual_seed(opt.seed)
    np.random.seed(opt.seed)
    
    # 设置设备
    if opt.device and opt.device.lower() != 'cpu':
        os.environ['CUDA_VISIBLE_DEVICES'] = opt.device
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    else:
        device = torch.device('cpu')
    
    print(f"使用设备: {device}")
    
    # 检查数据集
    data_dict = check_dataset(opt.data)
    
    # 创建保存目录
    save_dir = increment_path(Path(opt.project) / opt.name, 
                             exist_ok=opt.exist_ok)
    save_dir.mkdir(parents=True, exist_ok=True)
    
    # 保存训练配置
    with open(save_dir / 'opt.yaml', 'w') as f:
        yaml.dump(vars(opt), f, sort_keys=False)
    
    return device, data_dict, save_dir

def main(opt):
    """主训练函数"""
    print(colorstr('bright_green', 'bold', '=' * 50))
    print(colorstr('bright_green', 'bold', '吸烟检测系统训练开始'))
    print(colorstr('bright_green', 'bold', f"模型版本: {opt.model.upper()}"))
    print(colorstr('bright_green', 'bold', '=' * 50))
    
    # 设置环境
    device, data_dict, save_dir = setup_environment(opt)
    
    # 根据选择的模型调用对应的训练函数
    if opt.model == 'yolov5':
        train_yolov5(
            data=opt.data,
            epochs=opt.epochs,
            batch_size=opt.batch_size,
            imgsz=opt.img_size[0],
            device=device,
            weights=opt.weights,
            project=opt.project,
            name=opt.name,
            exist_ok=opt.exist_ok,
            workers=opt.workers,
            hyp=opt.hyp,
            patience=opt.patience,
            save_period=opt.save_period
        )
    elif opt.model == 'yolov6':
        train_yolov6(
            data=opt.data,
            epochs=opt.epochs,
            batch_size=opt.batch_size,
            img_size=opt.img_size,
            device=device,
            weights=opt.weights,
            project=opt.project,
            name=opt.name,
            exist_ok=opt.exist_ok,
            workers=opt.workers
        )
    elif opt.model == 'yolov7':
        train_yolov7(
            data=opt.data,
            epochs=opt.epochs,
            batch_size=opt.batch_size,
            img_size=opt.img_size[0],
            device=device,
            weights=opt.weights,
            project=opt.project,
            name=opt.name,
            exist_ok=opt.exist_ok,
            workers=opt.workers
        )
    elif opt.model == 'yolov8':
        train_yolov8(
            data=opt.data,
            epochs=opt.epochs,
            batch_size=opt.batch_size,
            imgsz=opt.img_size[0],
            device=device,
            model=opt.weights,
            project=opt.project,
            name=opt.name,
            exist_ok=opt.exist_ok,
            workers=opt.workers,
            patience=opt.patience
        )
    else:
        raise ValueError(f"不支持的模型版本: {opt.model}")
    
    print(colorstr('bright_green', 'bold', '=' * 50))
    print(colorstr('bright_green', 'bold', '训练完成!'))
    print(colorstr('bright_green', 'bold', f"结果保存在: {save_dir}"))
    print(colorstr('bright_green', 'bold', '=' * 50))

if __name__ == '__main__':
    opt = parse_opt()
    main(opt)
4.2.4 检测推理代码

python

"""
基于YOLO的吸烟检测系统推理脚本
支持实时视频、图片、摄像头输入
"""

import os
import sys
import argparse
import time
from pathlib import Path
import cv2
import torch
import numpy as np

# 添加项目根目录到路径
FILE = Path(__file__).resolve()
ROOT = FILE.parents[0]
if str(ROOT) not in sys.path:
    sys.path.append(str(ROOT))

from utils.general import (check_img_size, non_max_suppression, scale_boxes, 
                          increment_path, colorstr, check_requirements)
from utils.plots import Annotator, colors
from utils.torch_utils import select_device

class SmokingDetector:
    """吸烟检测器类"""
    
    def __init__(self, weights_path, model_version='yolov8', device='', img_size=640, conf_thres=0.25, iou_thres=0.45):
        """
        初始化检测器
        
        Args:
            weights_path: 模型权重路径
            model_version: YOLO版本 ['yolov5', 'yolov6', 'yolov7', 'yolov8']
            device: 运行设备
            img_size: 输入图像尺寸
            conf_thres: 置信度阈值
            iou_thres: IOU阈值
        """
        self.model_version = model_version
        self.img_size = img_size
        self.conf_thres = conf_thres
        self.iou_thres = iou_thres
        
        # 选择设备
        self.device = select_device(device)
        
        # 加载模型
        self.model = self.load_model(weights_path)
        
        # 获取类别名称
        self.names = self.model.module.names if hasattr(self.model, 'module') else self.model.names
        
        # 类别颜色
        self.colors = [[np.random.randint(0, 255) for _ in range(3)] for _ in range(len(self.names))]
        
        print(f"模型加载完成: {weights_path}")
        print(f"设备: {self.device}")
        print(f"类别: {self.names}")
    
    def load_model(self, weights_path):
        """加载模型"""
        check_requirements()
        
        if self.model_version == 'yolov5':
            from models.yolov5.models.experimental import attempt_load
            model = attempt_load(weights_path, device=self.device)
        elif self.model_version == 'yolov6':
            from models.yolov6.models.yolo import Model
            model = Model(weights_path, device=self.device)
        elif self.model_version == 'yolov7':
            from models.yolov7.models.experimental import attempt_load
            model = attempt_load(weights_path, device=self.device)
        elif self.model_version == 'yolov8':
            from ultralytics import YOLO
            model = YOLO(weights_path)
            model.to(self.device)
        else:
            raise ValueError(f"不支持的模型版本: {self.model_version}")
        
        # 设置为评估模式
        if self.model_version != 'yolov8':
            model.eval()
        
        return model
    
    def preprocess(self, img):
        """预处理图像"""
        # 调整图像尺寸
        img_resized = cv2.resize(img, (self.img_size, self.img_size))
        
        # 转换为RGB
        img_rgb = cv2.cvtColor(img_resized, cv2.COLOR_BGR2RGB)
        
        # 归一化并转换通道顺序
        img_normalized = img_rgb / 255.0
        img_transposed = np.transpose(img_normalized, (2, 0, 1))
        
        # 添加批次维度
        img_tensor = torch.from_numpy(img_transposed).float().unsqueeze(0)
        
        return img_tensor.to(self.device), img_resized
    
    def detect(self, img):
        """
        检测图像中的吸烟行为
        
        Args:
            img: 输入图像 (numpy数组)
            
        Returns:
            detections: 检测结果列表 [x1, y1, x2, y2, conf, class]
            img_annotated: 标注后的图像
        """
        # 原始图像尺寸
        orig_h, orig_w = img.shape[:2]
        
        # 预处理
        img_tensor, img_resized = self.preprocess(img)
        
        # 推理
        with torch.no_grad():
            if self.model_version == 'yolov8':
                results = self.model(img_tensor, conf=self.conf_thres, iou=self.iou_thres)[0]
                if results is not None:
                    detections = results.cpu().numpy()
                    # 转换为统一格式 [x1, y1, x2, y2, conf, class]
                    if detections.shape[1] == 6:
                        detections = detections
            else:
                pred = self.model(img_tensor)[0]
                # NMS
                pred = non_max_suppression(pred, self.conf_thres, self.iou_thres)
                
                detections = []
                if pred[0] is not None:
                    # 缩放边界框到原始图像尺寸
                    det = pred[0].cpu().numpy()
                    det[:, :4] = scale_boxes(img_resized.shape, det[:, :4], img.shape).round()
                    detections = det
        
        # 可视化
        img_annotated = self.visualize(img.copy(), detections)
        
        return detections, img_annotated
    
    def visualize(self, img, detections):
        """可视化检测结果"""
        annotator = Annotator(img, line_width=2, example=str(self.names))
        
        if len(detections) > 0:
            for det in detections:
                if len(det) >= 6:
                    x1, y1, x2, y2, conf, cls = det[:6]
                    label = f'{self.names[int(cls)]} {conf:.2f}'
                    annotator.box_label([x1, y1, x2, y2], label, color=self.colors[int(cls)])
        
        return annotator.result()
    
    def detect_video(self, video_path, output_path=None, show=True):
        """
        检测视频中的吸烟行为
        
        Args:
            video_path: 视频文件路径或摄像头索引
            output_path: 输出视频路径
            show: 是否实时显示
            
        Returns:
            检测结果统计
        """
        # 打开视频
        if isinstance(video_path, int) or video_path.isdigit():
            cap = cv2.VideoCapture(int(video_path))
        else:
            cap = cv2.VideoCapture(video_path)
        
        if not cap.isOpened():
            print(f"无法打开视频: {video_path}")
            return
        
        # 获取视频属性
        fps = int(cap.get(cv2.CAP_PROP_FPS))
        width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
        height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
        
        # 创建视频写入器
        if output_path:
            fourcc = cv2.VideoWriter_fourcc(*'mp4v')
            out = cv2.VideoWriter(output_path, fourcc, fps, (width, height))
        
        # 统计信息
        stats = {
            'total_frames': 0,
            'detection_frames': 0,
            'total_smoking': 0,
            'frames_with_smoking': 0
        }
        
        print(f"开始视频检测...")
        print(f"视频尺寸: {width}x{height}, FPS: {fps}")
        
        # 逐帧处理
        while True:
            ret, frame = cap.read()
            if not ret:
                break
            
            stats['total_frames'] += 1
            
            # 检测
            detections, frame_annotated = self.detect(frame)
            
            # 更新统计
            if len(detections) > 0:
                stats['detection_frames'] += 1
                smoking_count = sum(1 for det in detections if len(det) >= 6 and self.names[int(det[5])] == 'smoking')
                if smoking_count > 0:
                    stats['total_smoking'] += smoking_count
                    stats['frames_with_smoking'] += 1
                    
                    # 在帧上显示警告
                    cv2.putText(frame_annotated, "WARNING: Smoking Detected!", 
                               (50, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, 
                               (0, 0, 255), 3)
            
            # 显示
            if show:
                cv2.imshow('Smoking Detection', frame_annotated)
                if cv2.waitKey(1) & 0xFF == ord('q'):
                    break
            
            # 保存
            if output_path:
                out.write(frame_annotated)
        
        # 清理
        cap.release()
        if output_path:
            out.release()
        if show:
            cv2.destroyAllWindows()
        
        # 打印统计
        print(f"\n检测统计:")
        print(f"总帧数: {stats['total_frames']}")
        print(f"检测到目标的帧数: {stats['detection_frames']}")
        print(f"吸烟行为总数: {stats['total_smoking']}")
        print(f"包含吸烟行为的帧数: {stats['frames_with_smoking']}")
        
        return stats

def main():
    """主函数"""
    parser = argparse.ArgumentParser(description='吸烟检测推理脚本')
    parser.add_argument('--weights', type=str, default='runs/train/exp/weights/best.pt', 
                       help='模型权重路径')
    parser.add_argument('--source', type=str, default='0', 
                       help='输入源: 图片/视频路径, 或摄像头索引')
    parser.add_argument('--model', type=str, default='yolov8', 
                       choices=['yolov5', 'yolov6', 'yolov7', 'yolov8'],
                       help='YOLO版本')
    parser.add_argument('--img-size', type=int, default=640, 
                       help='推理尺寸')
    parser.add_argument('--conf-thres', type=float, default=0.25, 
                       help='置信度阈值')
    parser.add_argument('--iou-thres', type=float, default=0.45, 
                       help='NMS IOU阈值')
    parser.add_argument('--device', default='', 
                       help='运行设备')
    parser.add_argument('--output', type=str, default='', 
                       help='输出路径')
    parser.add_argument('--show', action='store_true', 
                       help='显示结果')
    
    opt = parser.parse_args()
    
    # 创建检测器
    detector = SmokingDetector(
        weights_path=opt.weights,
        model_version=opt.model,
        device=opt.device,
        img_size=opt.img_size,
        conf_thres=opt.conf_thres,
        iou_thres=opt.iou_thres
    )
    
    # 检查输入源类型
    source = opt.source
    is_image = source.lower().endswith(('.jpg', '.jpeg', '.png', '.bmp', '.tiff'))
    is_video = source.lower().endswith(('.mp4', '.avi', '.mov', '.mkv'))
    
    # 执行检测
    if is_image:
        # 图片检测
        img = cv2.imread(source)
        if img is None:
            print(f"无法读取图片: {source}")
            return
        
        detections, img_annotated = detector.detect(img)
        
        # 保存结果
        if opt.output:
            cv2.imwrite(opt.output, img_annotated)
            print(f"结果保存至: {opt.output}")
        
        # 显示
        if opt.show:
            cv2.imshow('Smoking Detection', img_annotated)
            cv2.waitKey(0)
            cv2.destroyAllWindows()
    
    else:
        # 视频或摄像头检测
        output_path = opt.output if opt.output else 'output.mp4'
        detector.detect_video(
            video_path=source,
            output_path=output_path,
            show=opt.show
        )

if __name__ == '__main__':
    main()
4.2.5 PySide6图形界面

python

"""
吸烟检测系统图形界面
基于PySide6开发
"""

import sys
import os
from pathlib import Path
import cv2
import numpy as np
from datetime import datetime

from PySide6.QtWidgets import (QApplication, QMainWindow, QWidget, QVBoxLayout, 
                              QHBoxLayout, QPushButton, QLabel, QFileDialog,
                              QMessageBox, QComboBox, QSpinBox, QDoubleSpinBox,
                              QGroupBox, QTextEdit, QProgressBar, QTabWidget,
                              QListWidget, QCheckBox, QSlider)
from PySide6.QtCore import Qt, QTimer, Signal, QThread, QSize
from PySide6.QtGui import QImage, QPixmap, QFont, QIcon

from smoking_detector import SmokingDetector

class DetectionThread(QThread):
    """检测线程"""
    frame_processed = Signal(np.ndarray, list)
    detection_stats = Signal(dict)
    finished = Signal()
    
    def __init__(self, detector, source):
        super().__init__()
        self.detector = detector
        self.source = source
        self.running = False
        self.stats = {
            'total_frames': 0,
            'detection_frames': 0,
            'total_smoking': 0,
            'frames_with_smoking': 0
        }
    
    def run(self):
        """线程运行函数"""
        self.running = True
        
        # 打开视频源
        if isinstance(self.source, int) or str(self.source).isdigit():
            cap = cv2.VideoCapture(int(self.source))
        else:
            cap = cv2.VideoCapture(self.source)
        
        if not cap.isOpened():
            print(f"无法打开视频源: {self.source}")
            self.finished.emit()
            return
        
        while self.running:
            ret, frame = cap.read()
            if not ret:
                break
            
            self.stats['total_frames'] += 1
            
            # 检测
            detections, frame_annotated = self.detector.detect(frame)
            
            # 更新统计
            if len(detections) > 0:
                self.stats['detection_frames'] += 1
                smoking_count = sum(1 for det in detections 
                                  if len(det) >= 6 and 
                                  self.detector.names[int(det[5])] == 'smoking')
                if smoking_count > 0:
                    self.stats['total_smoking'] += smoking_count
                    self.stats['frames_with_smoking'] += 1
                    
                    # 添加警告文字
                    cv2.putText(frame_annotated, "WARNING: Smoking Detected!", 
                               (50, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, 
                               (0, 0, 255), 3)
            
            # 发送信号
            self.frame_processed.emit(frame_annotated, detections)
            self.detection_stats.emit(self.stats.copy())
            
            # 控制帧率
            self.msleep(30)  # 约30FPS
        
        # 清理
        cap.release()
        self.finished.emit()
    
    def stop(self):
        """停止线程"""
        self.running = False

class SmokingDetectionGUI(QMainWindow):
    """吸烟检测系统主界面"""
    
    def __init__(self):
        super().__init__()
        self.detector = None
        self.detection_thread = None
        self.current_source = None
        self.is_detecting = False
        
        self.init_ui()
        self.init_detector()
    
    def init_ui(self):
        """初始化用户界面"""
        self.setWindowTitle("基于YOLO的吸烟检测系统 v1.0")
        self.setGeometry(100, 100, 1200, 800)
        
        # 设置应用图标
        if os.path.exists("icon.png"):
            self.setWindowIcon(QIcon("icon.png"))
        
        # 中央部件
        central_widget = QWidget()
        self.setCentralWidget(central_widget)
        
        # 主布局
        main_layout = QHBoxLayout(central_widget)
        
        # 左侧控制面板
        control_panel = self.create_control_panel()
        main_layout.addWidget(control_panel, 1)
        
        # 右侧显示区域
        display_panel = self.create_display_panel()
        main_layout.addWidget(display_panel, 2)
    
    def create_control_panel(self):
        """创建控制面板"""
        panel = QWidget()
        layout = QVBoxLayout(panel)
        
        # 标题
        title = QLabel("吸烟检测控制系统")
        title.setFont(QFont("Arial", 16, QFont.Bold))
        title.setAlignment(Qt.AlignCenter)
        layout.addWidget(title)
        
        # 模型选择组
        model_group = QGroupBox("模型配置")
        model_layout = QVBoxLayout()
        
        # YOLO版本选择
        model_layout.addWidget(QLabel("YOLO版本:"))
        self.model_combo = QComboBox()
        self.model_combo.addItems(["YOLOv5", "YOLOv6", "YOLOv7", "YOLOv8"])
        self.model_combo.setCurrentText("YOLOv8")
        model_layout.addWidget(self.model_combo)
        
        # 模型权重选择
        model_layout.addWidget(QLabel("模型权重:"))
        self.weight_combo = QComboBox()
        self.weight_combo.addItems(self.find_weights())
        model_layout.addWidget(self.weight_combo)
        
        # 加载模型按钮
        self.load_model_btn = QPushButton("加载模型")
        self.load_model_btn.clicked.connect(self.load_model)
        model_layout.addWidget(self.load_model_btn)
        
        model_group.setLayout(model_layout)
        layout.addWidget(model_group)
        
        # 检测参数组
        param_group = QGroupBox("检测参数")
        param_layout = QVBoxLayout()
        
        # 置信度阈值
        param_layout.addWidget(QLabel("置信度阈值:"))
        self.conf_slider = QSlider(Qt.Horizontal)
        self.conf_slider.setRange(0, 100)
        self.conf_slider.setValue(25)
        param_layout.addWidget(self.conf_slider)
        
        # IOU阈值
        param_layout.addWidget(QLabel("IOU阈值:"))
        self.iou_slider = QSlider(Qt.Horizontal)
        self.iou_slider.setRange(0, 100)
        self.iou_slider.setValue(45)
        param_layout.addWidget(self.iou_slider)
        
        param_group.setLayout(param_layout)
        layout.addWidget(param_group)
        
        # 输入源选择组
        source_group = QGroupBox("输入源")
        source_layout = QVBoxLayout()
        
        # 摄像头选择
        source_layout.addWidget(QLabel("摄像头:"))
        self.camera_combo = QComboBox()
        self.camera_combo.addItems(["摄像头0", "摄像头1", "摄像头2", "摄像头3"])
        source_layout.addWidget(self.camera_combo)
        
        # 文件选择
        self.file_btn = QPushButton("选择文件")
        self.file_btn.clicked.connect(self.select_file)
        source_layout.addWidget(self.file_btn)
        
        source_group.setLayout(source_layout)
        layout.addWidget(source_group)
        
        # 控制按钮组
        control_group = QGroupBox("控制")
        control_layout = QVBoxLayout()
        
        self.start_btn = QPushButton("开始检测")
        self.start_btn.clicked.connect(self.start_detection)
        self.start_btn.setEnabled(False)
        control_layout.addWidget(self.start_btn)
        
        self.stop_btn = QPushButton("停止检测")
        self.stop_btn.clicked.connect(self.stop_detection)
        self.stop_btn.setEnabled(False)
        control_layout.addWidget(self.stop_btn)
        
        self.screenshot_btn = QPushButton("截图保存")
        self.screenshot_btn.clicked.connect(self.save_screenshot)
        self.screenshot_btn.setEnabled(False)
        control_layout.addWidget(self.screenshot_btn)
        
        control_group.setLayout(control_layout)
        layout.addWidget(control_group)
        
        # 统计信息组
        stats_group = QGroupBox("统计信息")
        stats_layout = QVBoxLayout()
        
        self.stats_text = QTextEdit()
        self.stats_text.setReadOnly(True)
        stats_layout.addWidget(self.stats_text)
        
        stats_group.setLayout(stats_layout)
        layout.addWidget(stats_group)
        
        # 进度条
        self.progress_bar = QProgressBar()
        layout.addWidget(self.progress_bar)
        
        # 添加弹性空间
        layout.addStretch()
        
        return panel
    
    def create_display_panel(self):
        """创建显示面板"""
        panel = QWidget()
        layout = QVBoxLayout(panel)
        
        # 视频显示标签
        self.video_label = QLabel()
        self.video_label.setAlignment(Qt.AlignCenter)
        self.video_label.setMinimumSize(800, 600)
        self.video_label.setStyleSheet("border: 2px solid gray;")
        layout.addWidget(self.video_label)
        
        # 检测结果列表
        result_group = QGroupBox("检测结果")
        result_layout = QVBoxLayout()
        
        self.result_list = QListWidget()
        result_layout.addWidget(self.result_list)
        
        # 显示选项
        options_layout = QHBoxLayout()
        self.show_smoking_check = QCheckBox("显示吸烟检测")
        self.show_smoking_check.setChecked(True)
        options_layout.addWidget(self.show_smoking_check)
        
        self.show_person_check = QCheckBox("显示人物检测")
        options_layout.addWidget(self.show_person_check)
        
        self.show_cigarette_check = QCheckBox("显示香烟检测")
        options_layout.addWidget(self.show_cigarette_check)
        
        result_layout.addLayout(options_layout)
        result_group.setLayout(result_layout)
        layout.addWidget(result_group)
        
        return panel
    
    def init_detector(self):
        """初始化检测器"""
        # 尝试加载默认模型
        weights = self.find_weights()
        if weights:
            try:
                self.detector = SmokingDetector(
                    weights_path=weights[0],
                    model_version='yolov8',
                    device='',
                    img_size=640
                )
                self.load_model_btn.setText("模型已加载")
                self.start_btn.setEnabled(True)
            except Exception as e:
                QMessageBox.warning(self, "警告", f"加载模型失败: {str(e)}")
    
    def find_weights(self):
        """查找可用的模型权重"""
        weights = []
        weight_files = ["best.pt", "last.pt", "yolov8n.pt", "yolov5s.pt"]
        
        for file in weight_files:
            if os.path.exists(file):
                weights.append(file)
        
        # 检查runs/train目录
        train_dir = Path("runs/train")
        if train_dir.exists():
            for exp_dir in train_dir.iterdir():
                weights_dir = exp_dir / "weights"
                if weights_dir.exists():
                    for weight_file in weights_dir.glob("*.pt"):
                        weights.append(str(weight_file))
        
        return weights
    
    def load_model(self):
        """加载模型"""
        weight_path = self.weight_combo.currentText()
        model_version = self.model_combo.currentText().lower()
        
        if not weight_path or not os.path.exists(weight_path):
            QMessageBox.warning(self, "警告", "请选择有效的模型权重文件")
            return
        
        try:
            self.detector = SmokingDetector(
                weights_path=weight_path,
                model_version=model_version,
                device='',
                img_size=640,
                conf_thres=self.conf_slider.value() / 100,
                iou_thres=self.iou_slider.value() / 100
            )
            QMessageBox.information(self, "成功", "模型加载成功!")
            self.start_btn.setEnabled(True)
        except Exception as e:
            QMessageBox.critical(self, "错误", f"加载模型失败: {str(e)}")
    
    def select_file(self):
        """选择输入文件"""
        file_path, _ = QFileDialog.getOpenFileName(
            self, "选择文件", "", 
            "媒体文件 (*.mp4 *.avi *.mov *.mkv *.jpg *.jpeg *.png *.bmp)"
        )
        
        if file_path:
            self.current_source = file_path
            self.start_btn.setEnabled(True)
            QMessageBox.information(self, "信息", f"已选择文件: {file_path}")
    
    def start_detection(self):
        """开始检测"""
        if self.detector is None:
            QMessageBox.warning(self, "警告", "请先加载模型!")
            return
        
        # 确定输入源
        if hasattr(self, 'current_source') and self.current_source:
            source = self.current_source
        else:
            camera_index = int(self.camera_combo.currentText()[-1])
            source = camera_index
        
        # 创建检测线程
        self.detection_thread = DetectionThread(self.detector, source)
        self.detection_thread.frame_processed.connect(self.update_frame)
        self.detection_thread.detection_stats.connect(self.update_stats)
        self.detection_thread.finished.connect(self.detection_finished)
        
        # 更新UI状态
        self.is_detecting = True
        self.start_btn.setEnabled(False)
        self.stop_btn.setEnabled(True)
        self.screenshot_btn.setEnabled(True)
        
        # 开始检测
        self.detection_thread.start()
    
    def stop_detection(self):
        """停止检测"""
        if self.detection_thread and self.is_detecting:
            self.detection_thread.stop()
            self.is_detecting = False
            self.start_btn.setEnabled(True)
            self.stop_btn.setEnabled(False)
    
    def detection_finished(self):
        """检测完成"""
        self.is_detecting = False
        self.start_btn.setEnabled(True)
        self.stop_btn.setEnabled(False)
        self.video_label.clear()
        QMessageBox.information(self, "信息", "检测完成!")
    
    def update_frame(self, frame, detections):
        """更新视频帧"""
        # 转换图像格式
        rgb_image = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
        h, w, ch = rgb_image.shape
        bytes_per_line = ch * w
        
        # 创建QImage
        qt_image = QImage(rgb_image.data, w, h, bytes_per_line, QImage.Format_RGB888)
        
        # 缩放图像以适应标签
        scaled_image = qt_image.scaled(
            self.video_label.size(), 
            Qt.KeepAspectRatio, 
            Qt.SmoothTransformation
        )
        
        # 显示图像
        self.video_label.setPixmap(QPixmap.fromImage(scaled_image))
        
        # 更新结果列表
        self.update_result_list(detections)
    
    def update_result_list(self, detections):
        """更新结果列表"""
        self.result_list.clear()
        
        for det in detections:
            if len(det) >= 6:
                x1, y1, x2, y2, conf, cls = det[:6]
                class_name = self.detector.names[int(cls)]
                
                # 根据选项过滤显示
                if class_name == 'smoking' and not self.show_smoking_check.isChecked():
                    continue
                elif class_name == 'person' and not self.show_person_check.isChecked():
                    continue
                elif class_name == 'cigarette' and not self.show_cigarette_check.isChecked():
                    continue
                
                # 添加到列表
                item_text = f"{class_name}: 置信度 {conf:.2f}, 位置 ({x1:.0f}, {y1:.0f}) - ({x2:.0f}, {y2:.0f})"
                self.result_list.addItem(item_text)
    
    def update_stats(self, stats):
        """更新统计信息"""
        stats_text = f"""
总帧数: {stats['total_frames']}
检测到目标的帧数: {stats['detection_frames']}
吸烟行为总数: {stats['total_smoking']}
包含吸烟行为的帧数: {stats['frames_with_smoking']}

检测率: {stats['detection_frames']/max(stats['total_frames'], 1)*100:.1f}%
吸烟检测率: {stats['frames_with_smoking']/max(stats['total_frames'], 1)*100:.1f}%
        """
        self.stats_text.setText(stats_text)
    
    def save_screenshot(self):
        """保存截图"""
        if not hasattr(self, 'current_frame') or self.current_frame is None:
            QMessageBox.warning(self, "警告", "没有可保存的图像!")
            return
        
        # 生成文件名
        timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
        file_path = f"screenshot_{timestamp}.jpg"
        
        # 保存图像
        cv2.imwrite(file_path, self.current_frame)
        QMessageBox.information(self, "成功", f"截图已保存: {file_path}")
    
    def closeEvent(self, event):
        """关闭事件"""
        if self.is_detecting:
            self.stop_detection()
            event.ignore()
        else:
            event.accept()

def main():
    """主函数"""
    app = QApplication(sys.argv)
    
    # 设置应用程序样式
    app.setStyle('Fusion')
    
    # 创建并显示主窗口
    window = SmokingDetectionGUI()
    window.show()
    
    sys.exit(app.exec())

if __name__ == '__main__':
    main()
4.2.6 训练配置文件

hyp.yaml 超参数配置文件:

yaml

# YOLOv5超参数
lr0: 0.01  # 初始学习率
lrf: 0.01  # 最终学习率因子
momentum: 0.937  # 动量
weight_decay: 0.0005  # 权重衰减
warmup_epochs: 3.0  # 预热轮数
warmup_momentum: 0.8  # 预热动量
warmup_bias_lr: 0.1  # 预热偏置学习率
box: 0.05  # 边界框损失权重
cls: 0.5  # 类别损失权重
cls_pw: 1.0  # 类别BCL正样本权重
obj: 1.0  # 目标存在损失权重
obj_pw: 1.0  # 目标BCL正样本权重
iou_t: 0.20  # IoU训练阈值
anchor_t: 4.0  # 锚框阈值
fl_gamma: 0.0  # 焦点损失gamma
hsv_h: 0.015  # 图像HSV-色调增强
hsv_s: 0.7  # 图像HSV-饱和度增强
hsv_v: 0.4  # 图像HSV-明度增强
degrees: 0.0  # 图像旋转
translate: 0.1  # 图像平移
scale: 0.5  # 图像缩放
shear: 0.0  # 图像剪切
perspective: 0.0  # 图像透视变换
flipud: 0.0  # 图像上下翻转
fliplr: 0.5  # 图像左右翻转
mosaic: 1.0  # 图像马赛克增强
mixup: 0.0  # 图像MixUp增强
copy_paste: 0.0  # 图像复制粘贴增强

更多推荐