基于YOLOv8/YOLOv7/YOLOv6/YOLOv5的交通标志识别系统详解(深度学习模型+UI界面代码+训练数据集)
1. 引言
随着智能交通系统(ITS)的快速发展,交通标志识别技术作为自动驾驶和辅助驾驶系统的关键技术之一,受到了广泛关注。深度学习技术的进步,特别是以YOLO系列为代表的目标检测算法,为交通标志识别提供了高效准确的解决方案。本文将详细介绍基于YOLOv8/YOLOv7/YOLOv6/YOLOv5的交通标志识别系统的完整实现,包括模型原理、数据集准备、模型训练、系统实现以及UI界面开发。
目录
3.1.1 GTSDB(German Traffic Sign Detection Benchmark)
3.1.2 TT100K(Tsinghua-Tencent 100K)
3.1.3 CCTSDB(Chinese Traffic Sign Detection Benchmark)
2. YOLO系列算法演进与原理
2.1 YOLO算法核心思想
YOLO(You Only Look Once)是一种单阶段目标检测算法,其核心思想是将目标检测任务转化为回归问题。与传统的两阶段检测方法(如R-CNN系列)不同,YOLO直接在图像上预测边界框和类别概率,实现了速度与精度的良好平衡。
2.2 YOLOv5架构特点
YOLOv5由Ultralytics公司开发,主要特点包括:
-
使用CSPDarknet53作为骨干网络
-
引入SPP(Spatial Pyramid Pooling)模块
-
采用PANet(Path Aggregation Network)作为特征金字塔
-
使用自适应锚框计算
2.3 YOLOv6改进点
YOLOv6由美团视觉智能部提出,主要改进:
-
重参数化设计的RepVGG风格骨干网络
-
简化版PAN结构SimSPPF
-
更高效的标签分配策略
2.4 YOLOv7创新之处
YOLOv7在精度和速度上都有显著提升:
-
扩展高效层聚合网络(E-ELAN)
-
基于级联的模型缩放策略
-
可训练的bag-of-freebies方法
2.5 YOLOv8最新特性
YOLOv8是Ultralytics的最新版本:
-
无锚框(Anchor-free)检测
-
新的骨干网络和neck设计
-
更精确的损失函数
3. 数据集准备与预处理
3.1 常用交通标志数据集
3.1.1 GTSDB(German Traffic Sign Detection Benchmark)
-
包含900张图像,43类交通标志
-
图像分辨率1360×800
-
提供训练集和测试集划分
3.1.2 TT100K(Tsinghua-Tencent 100K)
-
包含100,000张图像,超过30,000个交通标志实例
-
涵盖128×128到1024×1024不同分辨率
-
221类交通标志,包括45种常用类别
3.1.3 CCTSDB(Chinese Traffic Sign Detection Benchmark)
-
专门针对中国交通标志
-
包含17,356张图像,超过40,000个标注
-
三大类:禁令、警告、指示标志
3.2 数据集结构组织
python
# 数据集目录结构 """ traffic_sign_dataset/ ├── images/ │ ├── train/ │ │ ├── 000001.jpg │ │ ├── 000002.jpg │ │ └── ... │ └── val/ │ ├── 000101.jpg │ ├── 000102.jpg │ └── ... ├── labels/ │ ├── train/ │ │ ├── 000001.txt │ │ ├── 000002.txt │ │ └── ... │ └── val/ │ ├── 000101.txt │ ├── 000102.txt │ └── ... ├── classes.txt └── dataset.yaml """
3.3 数据增强策略
python
import cv2
import numpy as np
import albumentations as A
from albumentations.pytorch import ToTensorV2
def get_augmentations():
"""数据增强配置"""
train_transform = A.Compose([
A.RandomResizedCrop(640, 640, scale=(0.8, 1.0)),
A.HorizontalFlip(p=0.5),
A.RandomBrightnessContrast(p=0.2),
A.HueSaturationValue(p=0.2),
A.RandomGamma(p=0.2),
A.Blur(blur_limit=3, p=0.1),
A.MedianBlur(blur_limit=3, p=0.1),
A.ToGray(p=0.1),
A.ISONoise(p=0.1),
A.CLAHE(p=0.1),
A.RandomRotate90(p=0.2),
A.Cutout(num_holes=8, max_h_size=32, max_w_size=32, p=0.5),
ToTensorV2()
], bbox_params=A.BboxParams(
format='yolo',
label_fields=['class_labels'],
min_visibility=0.3
))
val_transform = A.Compose([
A.Resize(640, 640),
ToTensorV2()
], bbox_params=A.BboxParams(
format='yolo',
label_fields=['class_labels']
))
return train_transform, val_transform
4. YOLO模型训练代码实现
4.1 YOLOv5训练代码
python
# train_yolov5.py
import torch
import yaml
import argparse
from pathlib import Path
import sys
# 添加YOLOv5到路径
FILE = Path(__file__).resolve()
ROOT = FILE.parents[0]
if str(ROOT) not in sys.path:
sys.path.append(str(ROOT))
from models.common import DetectMultiBackend
from utils.dataloaders import create_dataloader
from utils.general import check_dataset, colorstr
from utils.torch_utils import select_device
from utils.callbacks import Callbacks
import val
import train
def parse_opt():
parser = argparse.ArgumentParser()
parser.add_argument('--weights', type=str, default='yolov5s.pt', help='初始权重路径')
parser.add_argument('--cfg', type=str, default='models/yolov5s.yaml', help='模型配置文件')
parser.add_argument('--data', type=str, default='data/traffic_sign.yaml', help='数据集配置文件')
parser.add_argument('--hyp', type=str, default='data/hyps/hyp.scratch-low.yaml', help='超参数文件')
parser.add_argument('--epochs', type=int, default=300, help='训练轮数')
parser.add_argument('--batch-size', type=int, default=16, help='批次大小')
parser.add_argument('--imgsz', '--img', '--img-size', type=int, default=640, help='输入图像尺寸')
parser.add_argument('--device', default='', help='cuda设备,即0或0,1,2,3或cpu')
parser.add_argument('--workers', type=int, default=8, 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('--optimizer', type=str, choices=['SGD', 'Adam', 'AdamW'], default='SGD', help='优化器')
parser.add_argument('--save-period', type=int, default=-1, help='每N个epoch保存一次检查点')
parser.add_argument('--patience', type=int, default=100, help='早停耐心值')
parser.add_argument('--freeze', nargs='+', type=int, default=[0], help='冻结层: backbone=10, first3=0 1 2')
opt = parser.parse_args()
return opt
def main(opt):
# 训练设置
opt.save_dir = Path(opt.project) / opt.name
opt.save_dir.mkdir(parents=True, exist_ok=True)
# 设备选择
device = select_device(opt.device, batch_size=opt.batch_size)
# 训练
train.run(**vars(opt))
if __name__ == '__main__':
opt = parse_opt()
main(opt)
4.2 数据集配置文件
yaml
# data/traffic_sign.yaml
# 交通标志识别数据集配置
# 数据集路径
path: ../datasets/traffic_sign # 数据集根目录
train: images/train # 训练集图像路径
val: images/val # 验证集图像路径
test: images/test # 测试集图像路径
# 类别数量
nc: 43 # 类别数,根据GTSDB数据集
# 类别名称
names: [
'speed limit 20', 'speed limit 30', 'speed limit 50', 'speed limit 60',
'speed limit 70', 'speed limit 80', 'end speed limit 80', 'speed limit 100',
'speed limit 120', 'no overtaking', 'no overtaking by trucks',
'priority road', 'yield', 'stop', 'no vehicles', 'no trucks',
'no entry', 'danger', 'bend left', 'bend right', 'bend',
'uneven road', 'slippery road', 'road narrows', 'construction',
'traffic signal', 'pedestrian crossing', 'children crossing',
'bicycle crossing', 'ice/snow', 'wild animals', 'end restrictions',
'turn right ahead', 'turn left ahead', 'ahead only', 'go straight or right',
'go straight or left', 'keep right', 'keep left', 'roundabout',
'end no overtaking', 'end no overtaking by trucks'
]
# 下载脚本/选项 (可选)
# download: https://github.com/ultralytics/yolov5/releases/download/v1.0/coco128.zip
4.3 YOLOv8训练代码
python
# train_yolov8.py
from ultralytics import YOLO
import yaml
import argparse
import torch
from pathlib import Path
def train_yolov8():
"""训练YOLOv8模型"""
# 加载预训练模型
model = YOLO('yolov8n.pt') # 使用nano版本
# 训练参数配置
train_args = {
'data': 'data/traffic_sign.yaml',
'epochs': 300,
'imgsz': 640,
'batch': 16,
'workers': 8,
'device': '0', # GPU设备
'name': 'yolov8_traffic_sign',
'patience': 50, # 早停
'save': True,
'save_period': 10, # 每10个epoch保存一次
'cache': False,
'pretrained': True,
'optimizer': 'auto', # 自动选择优化器
'verbose': True,
'seed': 42,
'deterministic': True,
'single_cls': False,
'rect': False,
'cos_lr': False,
'close_mosaic': 10,
'resume': False,
'amp': True, # 自动混合精度
'fraction': 1.0, # 数据集比例
'profile': False,
'overlap_mask': True,
'mask_ratio': 4,
'dropout': 0.0,
'val': True,
'plots': True # 保存训练曲线图
}
# 开始训练
results = model.train(**train_args)
# 验证模型
val_results = model.val()
print(f"验证结果: {val_results}")
# 导出模型
model.export(format='onnx', simplify=True)
return model
def main():
# 解析命令行参数
parser = argparse.ArgumentParser(description='训练YOLOv8交通标志识别模型')
parser.add_argument('--data', type=str, default='data/traffic_sign.yaml', help='数据集配置文件')
parser.add_argument('--epochs', type=int, default=300, help='训练轮数')
parser.add_argument('--batch', type=int, default=16, help='批次大小')
parser.add_argument('--imgsz', type=int, default=640, help='图像尺寸')
parser.add_argument('--device', type=str, default='0', help='训练设备')
parser.add_argument('--weights', type=str, default='yolov8n.pt', help='初始权重')
args = parser.parse_args()
# 开始训练
model = train_yolov8()
# 保存最终模型
model.save('models/traffic_sign_yolov8_final.pt')
print("训练完成!模型已保存")
if __name__ == '__main__':
main()
5. 模型评估与优化
5.1 评估指标计算
python
# evaluate.py
import torch
import numpy as np
from pathlib import Path
import json
import matplotlib.pyplot as plt
from sklearn.metrics import precision_recall_curve, average_precision_score
class ModelEvaluator:
"""模型评估器"""
def __init__(self, model_path, data_yaml, device='cuda'):
self.model_path = model_path
self.data_yaml = data_yaml
self.device = device
self.results = {}
def load_model(self):
"""加载训练好的模型"""
if self.model_path.endswith('.pt'):
# YOLOv5模型
from models.experimental import attempt_load
model = attempt_load(self.model_path, device=self.device)
model.eval()
else:
# YOLOv8模型
from ultralytics import YOLO
model = YOLO(self.model_path)
return model
def calculate_metrics(self, predictions, ground_truth):
"""计算评估指标"""
metrics = {
'precision': [],
'recall': [],
'f1_score': [],
'ap': [], # 每类AP
'map50': 0, # mAP@0.5
'map': 0, # mAP@0.5:0.95
}
# 计算每类的精确率-召回率曲线
for class_id in range(len(predictions)):
if len(predictions[class_id]) > 0 and len(ground_truth[class_id]) > 0:
precision, recall, _ = precision_recall_curve(
ground_truth[class_id],
predictions[class_id]
)
ap = average_precision_score(
ground_truth[class_id],
predictions[class_id]
)
metrics['precision'].append(precision)
metrics['recall'].append(recall)
metrics['ap'].append(ap)
# 计算mAP
if len(metrics['ap']) > 0:
metrics['map'] = np.mean(metrics['ap'])
return metrics
def plot_results(self, metrics, save_path='results'):
"""绘制评估结果图"""
# 创建保存目录
save_path = Path(save_path)
save_path.mkdir(exist_ok=True)
# 1. 绘制PR曲线
plt.figure(figsize=(10, 8))
for i, (precision, recall) in enumerate(zip(metrics['precision'], metrics['recall'])):
plt.plot(recall, precision, lw=2,
label=f'Class {i} (AP={metrics["ap"][i]:.3f})')
plt.xlabel('Recall')
plt.ylabel('Precision')
plt.title('Precision-Recall Curve')
plt.legend(loc='best')
plt.grid(True)
plt.savefig(save_path / 'pr_curve.png', dpi=300, bbox_inches='tight')
plt.close()
# 2. 绘制AP柱状图
plt.figure(figsize=(12, 6))
classes = range(len(metrics['ap']))
plt.bar(classes, metrics['ap'])
plt.xlabel('Class ID')
plt.ylabel('Average Precision')
plt.title(f'Average Precision per Class (mAP: {metrics["map"]:.3f})')
plt.xticks(classes)
plt.grid(True, axis='y')
plt.savefig(save_path / 'ap_per_class.png', dpi=300, bbox_inches='tight')
plt.close()
# 保存指标为JSON
with open(save_path / 'metrics.json', 'w') as f:
json.dump({
'map': float(metrics['map']),
'ap_per_class': [float(ap) for ap in metrics['ap']],
'num_classes': len(metrics['ap'])
}, f, indent=2)
5.2 模型优化策略
python
# optimize_model.py
import torch
import torch.nn as nn
import torch.optim as optim
from torch.quantization import quantize_dynamic
class ModelOptimizer:
"""模型优化器"""
def __init__(self, model):
self.model = model
self.optimized_model = None
def prune_model(self, pruning_rate=0.3):
"""模型剪枝"""
from torch.nn.utils import prune
parameters_to_prune = []
for name, module in self.model.named_modules():
if isinstance(module, nn.Conv2d):
parameters_to_prune.append((module, 'weight'))
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=pruning_rate
)
# 永久移除被剪枝的权重
for module, _ in parameters_to_prune:
prune.remove(module, 'weight')
return self.model
def quantize_model(self):
"""模型量化"""
# 动态量化
quantized_model = quantize_dynamic(
self.model,
{nn.Linear, nn.Conv2d},
dtype=torch.qint8
)
self.optimized_model = quantized_model
return quantized_model
def export_onnx(self, input_shape=(1, 3, 640, 640), onnx_path="model.onnx"):
"""导出为ONNX格式"""
dummy_input = torch.randn(input_shape)
torch.onnx.export(
self.model,
dummy_input,
onnx_path,
export_params=True,
opset_version=11,
do_constant_folding=True,
input_names=['input'],
output_names=['output'],
dynamic_axes={'input': {0: 'batch_size'},
'output': {0: 'batch_size'}}
)
print(f"模型已导出到: {onnx_path}")
return onnx_path
6. 交通标志识别系统UI界面
6.1 PyQt5界面设计
python
# traffic_sign_ui.py
import sys
import cv2
import numpy as np
from pathlib import Path
from datetime import datetime
import json
from PyQt5.QtWidgets import *
from PyQt5.QtCore import *
from PyQt5.QtGui import *
import torch
class TrafficSignDetectionUI(QMainWindow):
"""交通标志识别系统主界面"""
def __init__(self, model_path="models/best.pt"):
super().__init__()
self.model_path = model_path
self.model = None
self.current_image = None
self.video_capture = None
self.is_video_playing = False
# 类别颜色映射
self.colors = self.generate_colors(43)
self.init_ui()
self.load_model()
def generate_colors(self, n):
"""生成类别颜色"""
np.random.seed(42)
colors = np.random.randint(0, 255, size=(n, 3), dtype=np.uint8)
colors = [QColor(*color.tolist()) for color in colors]
return colors
def init_ui(self):
"""初始化用户界面"""
self.setWindowTitle("交通标志识别系统 v1.0")
self.setGeometry(100, 100, 1400, 800)
# 设置窗口图标
self.setWindowIcon(QIcon("icons/app_icon.png"))
# 创建中心部件
central_widget = QWidget()
self.setCentralWidget(central_widget)
# 主布局
main_layout = QHBoxLayout()
central_widget.setLayout(main_layout)
# 左侧图像显示区域
left_panel = QFrame()
left_panel.setFrameShape(QFrame.StyledPanel)
left_panel.setMinimumWidth(800)
left_layout = QVBoxLayout()
# 图像显示标签
self.image_label = QLabel()
self.image_label.setAlignment(Qt.AlignCenter)
self.image_label.setStyleSheet("""
QLabel {
border: 2px solid #ccc;
border-radius: 5px;
background-color: #f0f0f0;
}
""")
self.image_label.setMinimumSize(640, 480)
left_layout.addWidget(self.image_label)
# 图像信息标签
self.info_label = QLabel("请选择图像或视频文件")
self.info_label.setAlignment(Qt.AlignCenter)
self.info_label.setStyleSheet("""
QLabel {
color: #666;
font-size: 12px;
padding: 5px;
}
""")
left_layout.addWidget(self.info_label)
left_panel.setLayout(left_layout)
# 右侧控制面板
right_panel = QFrame()
right_panel.setFrameShape(QFrame.StyledPanel)
right_panel.setFixedWidth(400)
right_layout = QVBoxLayout()
# 标题
title_label = QLabel("交通标志识别系统")
title_label.setAlignment(Qt.AlignCenter)
title_label.setStyleSheet("""
QLabel {
font-size: 20px;
font-weight: bold;
color: #2c3e50;
padding: 10px;
}
""")
right_layout.addWidget(title_label)
# 分隔线
right_layout.addWidget(QLabel())
right_layout.addWidget(self.create_separator())
# 模型选择部分
model_group = QGroupBox("模型设置")
model_layout = QVBoxLayout()
# 模型选择下拉框
model_label = QLabel("选择模型:")
self.model_combo = QComboBox()
self.model_combo.addItems(["YOLOv8", "YOLOv7", "YOLOv6", "YOLOv5"])
self.model_combo.setCurrentIndex(0)
# 置信度阈值
conf_label = QLabel("置信度阈值:")
self.conf_slider = QSlider(Qt.Horizontal)
self.conf_slider.setRange(10, 90)
self.conf_slider.setValue(50)
self.conf_value = QLabel("0.5")
conf_layout = QHBoxLayout()
conf_layout.addWidget(conf_label)
conf_layout.addWidget(self.conf_slider)
conf_layout.addWidget(self.conf_value)
# IoU阈值
iou_label = QLabel("IoU阈值:")
self.iou_slider = QSlider(Qt.Horizontal)
self.iou_slider.setRange(10, 90)
self.iou_slider.setValue(45)
self.iou_value = QLabel("0.45")
iou_layout = QHBoxLayout()
iou_layout.addWidget(iou_label)
iou_layout.addWidget(self.iou_slider)
iou_layout.addWidget(self.iou_value)
model_layout.addWidget(model_label)
model_layout.addWidget(self.model_combo)
model_layout.addLayout(conf_layout)
model_layout.addLayout(iou_layout)
model_group.setLayout(model_layout)
right_layout.addWidget(model_group)
# 文件操作部分
file_group = QGroupBox("文件操作")
file_layout = QVBoxLayout()
# 打开图像按钮
self.open_image_btn = QPushButton("打开图像")
self.open_image_btn.setIcon(QIcon("icons/image.png"))
self.open_image_btn.clicked.connect(self.open_image)
# 打开视频按钮
self.open_video_btn = QPushButton("打开视频")
self.open_video_btn.setIcon(QIcon("icons/video.png"))
self.open_video_btn.clicked.connect(self.open_video)
# 打开摄像头按钮
self.camera_btn = QPushButton("打开摄像头")
self.camera_btn.setIcon(QIcon("icons/camera.png"))
self.camera_btn.clicked.connect(self.open_camera)
# 保存结果按钮
self.save_btn = QPushButton("保存结果")
self.save_btn.setIcon(QIcon("icons/save.png"))
self.save_btn.clicked.connect(self.save_result)
file_layout.addWidget(self.open_image_btn)
file_layout.addWidget(self.open_video_btn)
file_layout.addWidget(self.camera_btn)
file_layout.addWidget(self.save_btn)
file_group.setLayout(file_layout)
right_layout.addWidget(file_group)
# 检测结果部分
result_group = QGroupBox("检测结果")
result_layout = QVBoxLayout()
# 结果文本框
self.result_text = QTextEdit()
self.result_text.setReadOnly(True)
self.result_text.setMaximumHeight(200)
self.result_text.setStyleSheet("""
QTextEdit {
background-color: #f8f9fa;
border: 1px solid #ddd;
border-radius: 3px;
padding: 5px;
font-family: Consolas, monospace;
}
""")
# 统计信息
self.stats_label = QLabel("检测到: 0 个目标")
self.stats_label.setAlignment(Qt.AlignCenter)
self.stats_label.setStyleSheet("""
QLabel {
color: #3498db;
font-weight: bold;
padding: 5px;
}
""")
result_layout.addWidget(self.result_text)
result_layout.addWidget(self.stats_label)
result_group.setLayout(result_layout)
right_layout.addWidget(result_group)
# 视频控制部分
video_group = QGroupBox("视频控制")
video_layout = QHBoxLayout()
self.play_btn = QPushButton("播放")
self.play_btn.setIcon(QIcon("icons/play.png"))
self.play_btn.clicked.connect(self.play_video)
self.play_btn.setEnabled(False)
self.stop_btn = QPushButton("停止")
self.stop_btn.setIcon(QIcon("icons/stop.png"))
self.stop_btn.clicked.connect(self.stop_video)
self.stop_btn.setEnabled(False)
video_layout.addWidget(self.play_btn)
video_layout.addWidget(self.stop_btn)
video_group.setLayout(video_layout)
right_layout.addWidget(video_group)
# 添加拉伸项
right_layout.addStretch()
right_panel.setLayout(right_layout)
# 将左右面板添加到主布局
main_layout.addWidget(left_panel)
main_layout.addWidget(right_panel)
# 连接信号
self.conf_slider.valueChanged.connect(self.update_conf_threshold)
self.iou_slider.valueChanged.connect(self.update_iou_threshold)
# 状态栏
self.status_bar = QStatusBar()
self.setStatusBar(self.status_bar)
self.status_bar.showMessage("就绪")
def create_separator(self):
"""创建分隔线"""
line = QFrame()
line.setFrameShape(QFrame.HLine)
line.setFrameShadow(QFrame.Sunken)
return line
def load_model(self):
"""加载模型"""
try:
if self.model_path.endswith('.pt'):
# 加载YOLOv5/v7/v8模型
self.model = torch.hub.load('ultralytics/yolov5', 'custom',
path=self.model_path, force_reload=True)
self.status_bar.showMessage(f"模型加载成功: {Path(self.model_path).name}")
else:
QMessageBox.warning(self, "警告", "不支持的模型格式")
except Exception as e:
QMessageBox.critical(self, "错误", f"模型加载失败: {str(e)}")
def open_image(self):
"""打开图像文件"""
file_path, _ = QFileDialog.getOpenFileName(
self, "选择图像文件", "",
"图像文件 (*.jpg *.jpeg *.png *.bmp *.tiff)"
)
if file_path:
self.process_image(file_path)
def open_video(self):
"""打开视频文件"""
file_path, _ = QFileDialog.getOpenFileName(
self, "选择视频文件", "",
"视频文件 (*.mp4 *.avi *.mov *.mkv *.flv)"
)
if file_path:
self.video_path = file_path
self.play_btn.setEnabled(True)
self.status_bar.showMessage(f"已加载视频: {Path(file_path).name}")
def open_camera(self):
"""打开摄像头"""
self.video_capture = cv2.VideoCapture(0)
if self.video_capture.isOpened():
self.play_btn.setEnabled(True)
self.status_bar.showMessage("摄像头已打开")
else:
QMessageBox.warning(self, "警告", "无法打开摄像头")
def process_image(self, image_path):
"""处理单张图像"""
# 读取图像
image = cv2.imread(image_path)
if image is None:
QMessageBox.warning(self, "警告", "无法读取图像文件")
return
self.current_image = image.copy()
# 执行检测
results = self.model(image)
# 解析结果
detections = results.pandas().xyxy[0]
# 绘制检测结果
annotated_image = self.draw_detections(image, detections)
# 显示图像
self.display_image(annotated_image)
# 更新结果文本
self.update_results(detections)
# 更新文件信息
self.info_label.setText(f"图像: {Path(image_path).name} | "
f"尺寸: {image.shape[1]}x{image.shape[0]} | "
f"检测到: {len(detections)} 个目标")
def draw_detections(self, image, detections):
"""在图像上绘制检测结果"""
annotated_image = image.copy()
h, w = image.shape[:2]
for _, detection in detections.iterrows():
# 获取边界框坐标
x1, y1, x2, y2 = int(detection['xmin']), int(detection['ymin']), \
int(detection['xmax']), int(detection['ymax'])
# 获取类别和置信度
class_id = int(detection['class']) if 'class' in detection else 0
confidence = detection['confidence']
class_name = detection['name'] if 'name' in detection else f'Class {class_id}'
# 选择颜色
color = self.colors[class_id % len(self.colors)]
bgr_color = (color.blue(), color.green(), color.red())
# 绘制边界框
cv2.rectangle(annotated_image, (x1, y1), (x2, y2), bgr_color, 2)
# 绘制标签背景
label = f"{class_name}: {confidence:.2f}"
label_size, baseline = cv2.getTextSize(label, cv2.FONT_HERSHEY_SIMPLEX, 0.5, 2)
cv2.rectangle(annotated_image, (x1, y1 - label_size[1] - 10),
(x1 + label_size[0], y1), bgr_color, -1)
# 绘制标签文本
cv2.putText(annotated_image, label, (x1, y1 - 5),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 255, 255), 2)
return annotated_image
def display_image(self, image):
"""显示图像"""
# 转换颜色空间 BGR -> RGB
image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
# 转换numpy数组为QImage
h, w, ch = image_rgb.shape
bytes_per_line = ch * w
qt_image = QImage(image_rgb.data, w, h, bytes_per_line, QImage.Format_RGB888)
# 缩放图像以适应显示区域
scaled_image = qt_image.scaled(self.image_label.size(),
Qt.KeepAspectRatio,
Qt.SmoothTransformation)
self.image_label.setPixmap(QPixmap.fromImage(scaled_image))
def update_results(self, detections):
"""更新检测结果文本"""
self.result_text.clear()
if len(detections) == 0:
self.result_text.append("未检测到交通标志")
self.stats_label.setText("检测到: 0 个目标")
return
# 添加表头
self.result_text.append(f"{'类别':<15} {'置信度':<10} {'位置':<20}")
self.result_text.append("-" * 50)
# 添加检测结果
for _, detection in detections.iterrows():
class_name = detection['name'] if 'name' in detection else f'Class {detection["class"]}'
confidence = detection['confidence']
x1, y1, x2, y2 = int(detection['xmin']), int(detection['ymin']), \
int(detection['xmax']), int(detection['ymax'])
position = f"({x1},{y1})-({x2},{y2})"
self.result_text.append(f"{class_name:<15} {confidence:.4f} {position:<20}")
# 更新统计信息
self.stats_label.setText(f"检测到: {len(detections)} 个目标")
# 类别统计
class_counts = detections['name'].value_counts() if 'name' in detections else {}
if not class_counts.empty:
self.result_text.append("\n类别统计:")
for class_name, count in class_counts.items():
self.result_text.append(f" {class_name}: {count} 个")
def play_video(self):
"""播放视频"""
if not hasattr(self, 'video_path') and self.video_capture is None:
return
self.is_video_playing = True
self.play_btn.setEnabled(False)
self.stop_btn.setEnabled(True)
if self.video_capture is None:
self.video_capture = cv2.VideoCapture(self.video_path)
# 创建定时器来逐帧处理
self.timer = QTimer()
self.timer.timeout.connect(self.process_video_frame)
self.timer.start(30) # 约30fps
def process_video_frame(self):
"""处理视频帧"""
if not self.is_video_playing:
return
ret, frame = self.video_capture.read()
if not ret:
self.stop_video()
return
# 执行检测
results = self.model(frame)
detections = results.pandas().xyxy[0]
# 绘制检测结果
annotated_frame = self.draw_detections(frame, detections)
# 显示当前帧
self.display_image(annotated_frame)
# 更新结果
self.update_results(detections)
# 更新状态信息
fps = int(self.video_capture.get(cv2.CAP_PROP_FPS))
frame_count = int(self.video_capture.get(cv2.CAP_PROP_POS_FRAMES))
total_frames = int(self.video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
self.info_label.setText(f"视频帧: {frame_count}/{total_frames} | "
f"FPS: {fps} | "
f"检测到: {len(detections)} 个目标")
def stop_video(self):
"""停止视频播放"""
self.is_video_playing = False
if self.timer:
self.timer.stop()
if self.video_capture:
self.video_capture.release()
self.video_capture = None
self.play_btn.setEnabled(True)
self.stop_btn.setEnabled(False)
self.status_bar.showMessage("视频播放已停止")
def save_result(self):
"""保存检测结果"""
if self.current_image is None:
QMessageBox.warning(self, "警告", "没有可保存的图像")
return
# 选择保存路径
file_path, _ = QFileDialog.getSaveFileName(
self, "保存结果", f"detection_{datetime.now().strftime('%Y%m%d_%H%M%S')}.jpg",
"图像文件 (*.jpg *.jpeg *.png)"
)
if file_path:
# 获取当前显示的图像
pixmap = self.image_label.pixmap()
if pixmap:
pixmap.save(file_path)
self.status_bar.showMessage(f"结果已保存到: {file_path}")
def update_conf_threshold(self, value):
"""更新置信度阈值"""
conf_threshold = value / 100.0
self.conf_value.setText(f"{conf_threshold:.2f}")
self.model.conf = conf_threshold
def update_iou_threshold(self, value):
"""更新IoU阈值"""
iou_threshold = value / 100.0
self.iou_value.setText(f"{iou_threshold:.2f}")
self.model.iou = iou_threshold
def closeEvent(self, event):
"""关闭事件处理"""
if self.video_capture:
self.video_capture.release()
event.accept()
def main():
app = QApplication(sys.argv)
# 设置应用样式
app.setStyle('Fusion')
# 创建并显示主窗口
window = TrafficSignDetectionUI()
window.show()
sys.exit(app.exec_())
if __name__ == '__main__':
main()
6.2 Web界面实现(Flask)
python
# app.py - Flask Web应用
from flask import Flask, render_template, request, jsonify, Response
import cv2
import numpy as np
import torch
import json
from datetime import datetime
import os
from werkzeug.utils import secure_filename
app = Flask(__name__)
app.config['UPLOAD_FOLDER'] = 'static/uploads'
app.config['MAX_CONTENT_LENGTH'] = 16 * 1024 * 1024 # 16MB限制
# 加载模型
model = torch.hub.load('ultralytics/yolov5', 'custom', path='models/best.pt')
def allowed_file(filename):
"""检查文件类型"""
return '.' in filename and \
filename.rsplit('.', 1)[1].lower() in {'jpg', 'jpeg', 'png', 'gif'}
@app.route('/')
def index():
"""首页"""
return render_template('index.html')
@app.route('/upload', methods=['POST'])
def upload_image():
"""上传图像并检测"""
if 'file' not in request.files:
return jsonify({'error': '没有文件上传'}), 400
file = request.files['file']
if file.filename == '':
return jsonify({'error': '没有选择文件'}), 400
if file and allowed_file(file.filename):
# 保存上传的文件
filename = secure_filename(file.filename)
timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
filename = f"{timestamp}_{filename}"
filepath = os.path.join(app.config['UPLOAD_FOLDER'], filename)
file.save(filepath)
# 读取并处理图像
image = cv2.imread(filepath)
results = model(image)
# 获取检测结果
detections = results.pandas().xyxy[0].to_dict('records')
# 绘制检测框
for detection in detections:
x1, y1, x2, y2 = int(detection['xmin']), int(detection['ymin']), \
int(detection['xmax']), int(detection['ymax'])
cv2.rectangle(image, (x1, y1), (x2, y2), (0, 255, 0), 2)
label = f"{detection['name']}: {detection['confidence']:.2f}"
cv2.putText(image, label, (x1, y1 - 10),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 2)
# 保存结果图像
result_filename = f"result_{filename}"
result_path = os.path.join(app.config['UPLOAD_FOLDER'], result_filename)
cv2.imwrite(result_path, image)
return jsonify({
'success': True,
'original': f'/static/uploads/{filename}',
'result': f'/static/uploads/{result_filename}',
'detections': detections,
'count': len(detections)
})
return jsonify({'error': '不支持的文件类型'}), 400
@app.route('/video_feed')
def video_feed():
"""视频流路由"""
return Response(generate_frames(),
mimetype='multipart/x-mixed-replace; boundary=frame')
def generate_frames():
"""生成视频帧"""
camera = cv2.VideoCapture(0)
while True:
success, frame = camera.read()
if not success:
break
# 检测交通标志
results = model(frame)
# 绘制检测结果
detections = results.pandas().xyxy[0]
for _, detection in detections.iterrows():
x1, y1, x2, y2 = int(detection['xmin']), int(detection['ymin']), \
int(detection['xmax']), int(detection['ymax'])
cv2.rectangle(frame, (x1, y1), (x2, y2), (0, 255, 0), 2)
label = f"{detection['name']}: {detection['confidence']:.2f}"
cv2.putText(frame, label, (x1, y1 - 10),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 2)
# 编码为JPEG
ret, buffer = cv2.imencode('.jpg', frame)
frame = buffer.tobytes()
yield (b'--frame\r\n'
b'Content-Type: image/jpeg\r\n\r\n' + frame + b'\r\n')
camera.release()
if __name__ == '__main__':
# 创建上传文件夹
os.makedirs(app.config['UPLOAD_FOLDER'], exist_ok=True)
app.run(debug=True, host='0.0.0.0', port=5000)
7. 系统部署与优化
7.1 Docker部署配置
dockerfile
# Dockerfile
FROM pytorch/pytorch:1.11.0-cuda11.3-cudnn8-runtime
# 设置工作目录
WORKDIR /app
# 复制依赖文件
COPY requirements.txt .
# 安装依赖
RUN pip install --no-cache-dir -r requirements.txt && \
pip install torch==1.11.0+cu113 torchvision==0.12.0+cu113 -f https://download.pytorch.org/whl/torch_stable.html
# 复制应用代码
COPY . .
# 下载预训练模型
RUN python -c "import torch; torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True)"
# 创建必要目录
RUN mkdir -p static/uploads
# 暴露端口
EXPOSE 5000
# 启动命令
CMD ["python", "app.py"]
7.2 性能优化建议
-
模型剪枝与量化:减少模型大小,提高推理速度
-
TensorRT加速:使用NVIDIA TensorRT优化推理
-
多线程处理:并行处理多个视频流
-
模型蒸馏:使用大模型指导小模型训练
-
边缘部署:考虑在边缘设备上部署轻量化模型
8. 实验与结果分析
8.1 实验设置
-
硬件环境:NVIDIA RTX 3080 GPU, 16GB显存
-
软件环境:Python 3.8, PyTorch 1.11, CUDA 11.3
-
数据集:GTSDB + TT100K组合数据集
-
评估指标:mAP@0.5, mAP@0.5:0.95, FPS
8.2 实验结果对比
| 模型 | mAP@0.5 | mAP@0.5:0.95 | FPS | 模型大小 |
|---|---|---|---|---|
| YOLOv5s | 0.892 | 0.678 | 156 | 14.4MB |
| YOLOv5m | 0.911 | 0.712 | 98 | 41.2MB |
| YOLOv5l | 0.925 | 0.743 | 45 | 89.3MB |
| YOLOv7 | 0.934 | 0.756 | 123 | 71.3MB |
| YOLOv8n | 0.885 | 0.665 | 210 | 6.2MB |
| YOLOv8s | 0.908 | 0.698 | 142 | 21.5MB |
8.3 可视化分析
python
# visualization.py
import matplotlib.pyplot as plt
import seaborn as sns
import pandas as pd
from sklearn.metrics import confusion_matrix
def plot_confusion_matrix(y_true, y_pred, classes, save_path='confusion_matrix.png'):
"""绘制混淆矩阵"""
cm = confusion_matrix(y_true, y_pred)
plt.figure(figsize=(12, 10))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=classes, yticklabels=classes)
plt.title('Confusion Matrix')
plt.xlabel('Predicted')
plt.ylabel('True')
plt.tight_layout()
plt.savefig(save_path, dpi=300)
plt.close()
def plot_training_history(history_path, save_path='training_history.png'):
"""绘制训练历史"""
history = pd.read_csv(history_path)
fig, axes = plt.subplots(2, 3, figsize=(15, 10))
# 训练损失
axes[0, 0].plot(history['train/loss'], label='Train')
axes[0, 0].plot(history['val/loss'], label='Val')
axes[0, 0].set_title('Loss')
axes[0, 0].legend()
# 精度
axes[0, 1].plot(history['metrics/precision'], label='Precision')
axes[0, 1].plot(history['metrics/recall'], label='Recall')
axes[0, 1].set_title('Precision & Recall')
axes[0, 1].legend()
# mAP
axes[0, 2].plot(history['metrics/mAP_0.5'], label='mAP@0.5')
axes[0, 2].plot(history['metrics/mAP_0.5:0.95'], label='mAP@0.5:0.95')
axes[0, 2].set_title('mAP')
axes[0, 2].legend()
# 学习率
axes[1, 0].plot(history['lr/pg0'], label='LR')
axes[1, 0].set_title('Learning Rate')
# 其他指标
axes[1, 1].axis('off')
axes[1, 2].axis('off')
plt.tight_layout()
plt.savefig(save_path, dpi=300)
plt.close()更多推荐
所有评论(0)