–从零开始,手把手教你打造自己的AI火灾检测系统,告别命令行,点点鼠标就能训练模型!
#创作灵感#
火灾检测一直是计算机视觉领域的重要研究课题,最有希望不造成重大经济损失的阶段就是火灾的早早期阶段。我的研究立意在于创作一个简单的GUI界面,能够在拥有数据记得情况下直接在这个GUI平台上进行训练测试,以减少不必要的时间浪费。
我将GUI的完整代码放在文末,大家可以随意取用,因项目工程过大,没办法直接附上,大家有需要可以私信我或者在评论区留言,我看到后会一一回复。也欢迎有志同道合的通知找我交流。
1. 为什么要做这个项目?
1.1 一个让人心痛的现实
2023年,全球发生了超过3.7万起森林火灾,烧毁了超过1.5亿公顷的森林面积。这不仅仅是数字,而是无数动物的家园,是我们呼吸的空气,是地球的肺。
作为技术人员,我一直在想:我能做什么?
1.2 AI能做什么?
传统的火灾检测依赖于:
-
卫星遥感(时效性差)
-
瞭望塔人工观察(覆盖面小)
-
传感器网络(成本高)
而基于深度学习的计算机视觉技术,可以:
-
实时监测:秒级响应
-
覆盖广泛:只要有摄像头就能部署
-
成本低廉:现有摄像头+软件升级
-
准确率高:YOLOv8在火灾检测上可达90%+准确率
1.3 为什么需要GUI?
在我调研现有的火灾检测解决方案时,发现了一个严重问题:
绝大多数方案都是命令行式的!
# 训练一个模型需要记住这么多参数
yolo task=detect mode=train model=yolov8s.pt data=fire.yaml epochs=100 imgsz=640 batch=16 lr0=0.01 workers=8 device=0 amp=True plots=True …
这对非技术背景的消防人员、林业管理人员来说,简直是天书!
我的目标:让消防员都能用的AI工具!
先看一下整体的框架界面!

1.4 项目的核心价值
这个项目的意义不仅仅是技术实现,更是:
降低门槛:让不会编程的人也能用AI
提高效率:图形化界面比命令行快10倍
减少错误:参数填写有提示,不会记错
即开即用:一键启动,无需配置
开源共享:任何人都可以改进和使用
2. 项目背景与技术选型
2.1 为什么选择YOLOv8?
YOLO系列已经发展到第8代,选择它的原因:
| 速度 | ⭐⭐⭐⭐⭐ 实时检测 | 较慢 |
| 准确率 | ⭐⭐⭐⭐⭐ SOTA级别 | 参差不齐 |
| 易用性 | ⭐⭐⭐⭐⭐ 开箱即用 | 需要大量配置 |
| 社区支持 | ⭐⭐⭐⭐⭐ 最活跃 | 较少 |
| 预训练模型 | ⭐⭐⭐⭐⭐ 多种尺寸 | 选择少 |
YOLOv8的架构优势:
-
C2f模块:更高效的特征提取
-
Decoupled Head:分类和回归分离
-
Task Alignment:更准确的边界框预测
-
Mosaic Augmentation:数据增强,提升泛化能力
2.2 为什么选PyQt5做GUI?
| 界面美观度 | ⭐⭐⭐⭐⭐ | ⭐⭐⭐ | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐⭐ |
| 性能 | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐ | ⭐⭐⭐ | ⭐⭐⭐ |
| 开发效率 | ⭐⭐⭐⭐ | ⭐⭐⭐⭐⭐ | ⭐⭐⭐ | ⭐⭐⭐⭐ |
| 跨平台 | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐⭐ |
| 学习曲线 | ⭐⭐⭐ | ⭐⭐⭐⭐⭐ | ⭐⭐ | ⭐⭐ |
PyQt5在功能、美观和性能之间取得了最佳平衡。
2.3 完整技术栈
前端:PyQt5 + QSS自定义样式
├── 双主题支持(暗黑/明亮)
├── 自定义控件(图像查看器、日志控制台)
└── 响应式布局
后端:Ultralytics YOLOv8 + PyTorch
├── 模型训练管理
├── 推理引擎
└── 结果后处理
数据流:subprocess + threading
├── 实时输出捕获
├── 进度解析
└── 错误诊断
可视化:matplotlib + PIL
├── 训练曲线绘制
├── 混淆矩阵展示
└── 预测结果展示
3. 环境搭建:一步步配置开发环境
3.1 硬件要求
最低配置(CPU训练):
-
CPU: Intel Core i5 或 AMD 同等
-
内存: 8GB DDR4
-
硬盘: 20GB 可用空间
-
系统: Windows 10/11, Ubuntu 20.04+, macOS 10.15+
推荐配置(GPU训练):
-
CPU: Intel Core i7/i9 或 AMD Ryzen 7/9
-
内存: 16GB+
-
GPU: NVIDIA RTX 3060+ (6GB+ VRAM)
-
硬盘: 50GB+ SSD
-
CUDA: 11.8+ 和 cuDNN 8.7+
3.2 第一步:安装Python
我推荐使用 Anaconda 或 Miniconda 管理Python环境:
# 下载并安装 Miniconda
# https://docs.conda.io/en/latest/miniconda.html
# 创建专门的环境
conda create -n yolov8_fire python=3.10
# 激活环境
conda activate yolov8_fire
为什么用3.10而不是3.12?
-
PyTorch 对3.12的支持还不完美
-
许多库在3.10上经过充分测试
-
兼容性最好,bug最少
3.3 第二步:安装PyTorch
根据你的CUDA版本选择安装命令:
# CUDA 11.8 (推荐)
pip install torch==2.0.1 torchvision==0.15.2 –index-url https://download.pytorch.org/whl/cu118
# CUDA 12.1
pip install torch==2.1.0 torchvision==0.16.0 –index-url https://download.pytorch.org/whl/cu121
# CPU版本
pip install torch torchvision –index-url https://download.pytorch.org/whl/cpu
验证安装:
import torch
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"GPU型号: {torch.cuda.get_device_name(0)}")
print(f"显存: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB")
3.4 第三步:安装Ultralytics YOLOv8
# 安装最新版本
pip install ultralytics
# 或者安装特定版本(更稳定)
pip install ultralytics==8.0.20
验证安装:
import ultralytics
print(f"Ultralytics版本: {ultralytics.__version__}")
3.5 第四步:安装GUI相关依赖
# PyQt5
pip install PyQt5 PyQtWebEngine
# 图像处理
pip install opencv-python pillow
# 数据可视化
pip install matplotlib
# 其他工具
pip install pyyaml numpy
3.6 第五步:一次性安装所有依赖
创建 requirements.txt:
torch==2.0.1
torchvision==0.15.2
ultralytics==8.0.20
PyQt5==5.15.9
opencv-python==4.8.1.78
pillow==10.1.0
matplotlib==3.7.3
pyyaml==6.0.1
numpy==1.24.3
一键安装:
pip install -r requirements.txt
3.7 第六步:验证完整环境
创建 test_env.py:
import sys
import torch
import ultralytics
import PyQt5
import cv2
from PIL import Image
import matplotlib
import yaml
import numpy as np
def test_environment():
print("=" * 50)
print("环境测试报告")
print("=" * 50)
print(f"Python版本: {sys.version}")
print(f"PyTorch: {torch.__version__}")
print(f"Ultralytics: {ultralytics.__version__}")
print(f"PyQt5: {PyQt5.QtCore.QT_VERSION_STR}")
print(f"OpenCV: {cv2.__version__}")
print(f"Pillow: {Image.__version__}")
print(f"Matplotlib: {matplotlib.__version__}")
print(f"NumPy: {np.__version__}")
if torch.cuda.is_available():
print(f"✅ GPU: {torch.cuda.get_device_name(0)}")
else:
print("⚠️ GPU不可用,将使用CPU训练")
print("✅ 环境配置完成!")
if __name__ == "__main__":
test_environment()
运行:
python test_env.py
4. 数据集准备:从零开始构建火灾数据集
4.1 数据集的重要性
"Garbage in, garbage out" – 这句话在AI领域尤其正确
一个好的数据集需要:
-
多样性:不同场景、光照、角度
-
准确性:标注精确,类别明确
-
数量:每个类别至少1000+样本
-
平衡:各类别样本数量均衡


4.2 数据来源
1. Roboflow Universe(推荐)
-
全球最大的开源数据集平台
-
搜索 "fire detection" 即可找到多个数据集
-
一键导出YOLOv8格式
2. 公开数据集
| FLAME | ~50,000张 | 森林火灾 | GitHub |
| FIRE | ~2,000张 | 室内火灾 | Kaggle |
| DeepFire | ~1,500张 | 多种火灾场景 | GitHub |
3. 自己采集
-
从视频中提取帧
-
网络爬虫爬取
-
实地拍摄
4. 数据增强
-
旋转、缩放、翻转
-
色彩变换(亮度、对比度、饱和度)
-
添加噪声
-
Mosaic增强
4.3 数据集结构(YOLOv8格式)
datasets/fire-8/
├── train/
│ ├── images/
│ │ ├── image_001.jpg
│ │ ├── image_002.jpg
│ │ └── …
│ └── labels/
│ ├── image_001.txt
│ ├── image_002.txt
│ └── …
├── valid/
│ ├── images/
│ └── labels/
├── test/
│ ├── images/
│ └── labels/
└── data.yaml
4.4 标签文件格式(YOLO格式)
每个 image_xxx.txt 文件包含多行,每行一个目标:
class_id x_center y_center width height
其中坐标是归一化的(0-1之间):
-
x_center:中心点x坐标 / 图片宽度
-
y_center:中心点y坐标 / 图片高度
-
width:目标宽度 / 图片宽度
-
height:目标高度 / 图片高度
示例:
0 0.45 0.55 0.30 0.40
1 0.78 0.32 0.15 0.20
4.5 data.yaml配置文件
# 数据集路径(相对于此文件的位置)
path: . # 当前目录
# 训练/验证/测试集路径(相对于 path)
train: train/images
val: valid/images
test: test/images
# 类别数量
nc: 2
# 类别名称(与标签文件中的类别索引对应)
names:
0: fire
1: smoke
4.6 数据预处理最佳实践
图片尺寸统一:建议640×640或800×800
格式统一:全部转为JPG(压缩)或PNG(无损)
去除重复:使用哈希值去重
质量过滤:删除模糊、过曝的图片
标注检查:确保标签与图片对应
4.7 数据集划分策略
总数据集: 100%
├── 训练集: 70%
├── 验证集: 20%
└── 测试集: 10%
交叉验证:
-
K-Fold: 5折交叉验证
-
时间序列划分:按时间顺序
5. GUI界面设计:让AI训练像玩游戏一样简单
5.1 UI/UX设计原则
在设计这个GUI时,我遵循了几个核心原则:
一目了然:用户打开程序就知道能做什么
流程清晰:按照"准备数据 → 配置参数 → 训练 → 推理"的顺序
反馈即时:每个操作都有明确的反馈
容错性高:即使操作错误也不会崩溃
美观大方:科技感十足,让人愿意使用
5.2 配色方案
暗黑主题(默认):
背景色: #1a1a2e (深蓝黑)
卡片背景: #16213e (深蓝)
强调色: #e94560 (活力红)
次级色: #0f3460 (深蓝)
文字色: #e0e0e0 (灰白)
边框色: #2d2d4a (深紫灰)
明亮主题:
背景色: #f5f5f5 (浅灰)
卡片背景: #ffffff (纯白)
强调色: #e94560 (活力红)
文字色: #333333 (深灰)
边框色: #d0d0d0 (浅灰)
5.3 布局设计
+————————————————–+
| 🔥 YOLOv8 Early Fire Detection |
| 文件 视图 帮助 |
+————————————————–+
| 🌓 主题 | 🚀 训练 | 🔍 推理 | 🗑️ 清空 | 🐛 调试 |
+————————————————–+
| |
| +———————————————-+ |
| | 🏠 欢迎 🚀 训练 🔍 推理 📈 可视化 📟 控制台| |
| +———————————————-+ |
| |
| [主内容区域] |
| |
| +———————————————-+ |
| | 状态栏: Ready [进度条] | |
| +———————————————-+ |
+————————————————–+

5.4 五个核心标签页
🏠 欢迎页
-
项目标题和标语
-
三个功能卡片(训练、推理、可视化)
-
快速入门指南
🚀 训练页
-
左侧:配置面板(数据集、模型、超参数)
-
右侧:训练结果预览(图片查看器)
🔍 推理页
-
左侧:配置面板(模型、推理源)
-
右侧:推理结果预览(图片查看器 + 信息)
📈 可视化页
-
下拉选择器(训练结果、混淆矩阵、验证预测、自定义图片)
-
Matplotlib绘图区域
-
状态信息
📟 控制台页
-
实时日志输出
-
清空和保存功能
-
时间戳标记
5.5 自定义控件
1. ConsoleOutput(控制台)
class ConsoleOutput(QPlainTextEdit):
def append(self, text):
# 带时间戳的输出
timestamp = datetime.now().strftime("%H:%M:%S")
self.appendPlainText(f"[{timestamp}] {text}")
# 自动滚动到底部
self.verticalScrollBar().setValue(self.verticalScrollBar().maximum())
2. ImageViewer(图像查看器)
class ImageViewer(QWidget):
# 支持缩放、平移、自适应
def zoom_in(self): …
def zoom_out(self): …
def fit_view(self): …
3. TrainingWorker(训练线程)
class TrainingWorker(QThread):
# 多线程训练,不阻塞UI
output_signal = pyqtSignal(str)
finished_signal = pyqtSignal(int, str)
progress_signal = pyqtSignal(int)
6. 核心功能实现:训练、推理、可视化
6.1 训练功能实现
命令构建:
def start_training(self):
# 获取用户配置
model = self.model_size.currentText().split()[0]
data = self.data_yaml_path.text()
epochs = self.epochs.value()
imgsz = self.img_size.value()
batch = self.batch_size.value()
lr = self.lr.value()
workers = self.workers.value()
device = self.device.currentText().split()[0]
# 构建YOLO命令
cmd = f"yolo task=detect mode=train model={model} data={data} epochs={epochs} imgsz={imgsz} batch={batch} lr0={lr} workers={workers} device={device} amp=False"
# 添加可选参数
if self.pretrained_check.isChecked():
cmd += " pretrained=True"
cmd += " plots=True"
# 启动训练线程
self.training_worker = TrainingWorker(cmd, os.getcwd())
self.training_worker.output_signal.connect(self.console.append)
self.training_worker.progress_signal.connect(self.progress_bar.setValue)
self.training_worker.finished_signal.connect(self.training_finished)
self.training_worker.start()
实时输出解析:
def read_stdout():
for line in self.process.stdout:
# 解析训练进度
if 'epoch' in line.lower():
match = re.search(r'epoch[:\\s]+(\\d+)/(\\d+)', line, re.IGNORECASE)
if match:
current = int(match.group(1))
total = int(match.group(2))
progress = int((current / total) * 100)
self.progress_signal.emit(progress)
6.2 推理功能实现
def run_inference(self):
# 获取配置
model = self.model_weights.text()
source = self.source_path.text()
conf = self.confidence.value()
save = self.save_results_check.isChecked()
# 构建推理命令
cmd = f"yolo task=detect mode=predict model={model} conf={conf} source={source} save={str(save).lower()}"
# 运行推理
infer_worker = TrainingWorker(cmd, os.getcwd())
infer_worker.finished_signal.connect(self.inference_finished)
infer_worker.start()
6.3 可视化功能实现
def update_visualization(self):
selection = self.viz_combo.currentText()
if selection == "Training Results":
path = os.path.join(TRAIN_RESULTS, "results.png")
self.viz_widget.plot_image(path)
self.viz_info.setText("📊 Training Results Overview")
elif selection == "Confusion Matrix":
path = os.path.join(TRAIN_RESULTS, "confusion_matrix.png")
self.viz_widget.plot_image(path)
self.viz_info.setText("📊 Confusion Matrix")
elif selection == "Validation Predictions":
path = os.path.join(TRAIN_RESULTS, "val_batch0_pred.jpg")
self.viz_widget.plot_image(path)
self.viz_info.setText("🔍 Validation Batch Predictions")
6.4 错误诊断系统
这是我最得意的功能之一:
def analyze_error(self, error_output):
"""智能分析错误并给出建议"""
error_text = ' '.join(error_output).lower()
if 'cuda' in error_text or 'cudnn' in error_text:
return "💡 CUDA相关错误,建议:\\n1. 设置 device=cpu\\n2. 检查CUDA版本\\n3. 减少batch size"
if 'out of memory' in error_text or 'oom' in error_text:
return "💡 内存不足,建议:\\n1. 减少batch size\\n2. 减少image size\\n3. 使用CPU训练"
if 'file not found' in error_text:
return "💡 文件不存在,检查:\\n1. data.yaml中的路径\\n2. 图片和标签是否存在\\n3. 路径格式是否正确"
if 'attribute' in error_text or 'has no' in error_text:
return "💡 版本兼容性问题,建议:\\n1. 更新ultralytics: pip install -U ultralytics\\n2. 降级到稳定版本"
return "💡 未知错误,请查看详细日志"
7. 踩坑实录:我遇到的10个坑及解决方案
坑1:模型文件损坏
错误信息:
RuntimeError: PytorchStreamReader failed reading zip archive:
failed finding central directory
解决过程:
检查文件大小:发现只有几KB
删除重新下载
下载完成后再检查大小
根本原因:
-
下载过程中网络中断
-
磁盘空间不足
-
权限问题
永久解决方案:
def validate_model(model_path):
if os.path.exists(model_path):
size = os.path.getsize(model_path)
if size < 1000000: # 小于1MB
os.remove(model_path)
print(f"⚠️ {model_path} 损坏,已删除")
return False
return True
坑2:AMP检查失败
错误信息:
assert amp_allclose(YOLO("yolo26n.pt"), im)
原因分析:
-
Ultralytics 8.4.x在AMP检查时加载了不存在的模型
-
这是开发版bug
解决方案:
# 在训练命令中添加 amp=False
cmd += " amp=False"
坑3:数据集路径错误
错误信息:
FileNotFoundError: data.yaml not found
解决过程:
检查data.yaml是否存在
检查路径是否正确
使用绝对路径
永久解决方案:
# 自动检测并设置路径
DATASET_PATH = os.path.join(PROJECT_ROOT, "datasets", "fire-8")
DATA_YAML = os.path.join(DATASET_PATH, "data.yaml")
if os.path.exists(DATA_YAML):
self.data_yaml_path.setText(DATA_YAML)
else:
self.console.append_warning(f"data.yaml not found at {DATA_YAML}")
坑4:GPU内存不足
错误信息:
CUDA out of memory
解决方案:
减小batch size(16 → 8 → 4)
减小image size(640 → 416)
使用混合精度训练
使用梯度累积
坑5:标签文件格式错误
问题: 标签文件中包含超出范围(0-1)的坐标值
解决方案:
def fix_labels(label_dir):
for file in glob.glob(os.path.join(label_dir, "*.txt")):
with open(file, 'r') as f:
lines = f.readlines()
fixed_lines = []
for line in lines:
parts = line.strip().split()
if len(parts) == 5:
# 确保坐标在0-1之间
cls_id = int(parts[0])
coords = [float(x) for x in parts[1:5]]
coords = [max(0, min(1, x)) for x in coords]
fixed_lines.append(f"{cls_id} {' '.join([f'{x:.6f}' for x in coords])}")
with open(file, 'w') as f:
f.write('\\n'.join(fixed_lines))
坑6:类别数量不匹配
问题: data.yaml中的nc与实际类别数不符
解决方案:
def check_classes(data_yaml, label_dirs):
# 统计实际类别数
actual_classes = set()
for label_dir in label_dirs:
for file in glob.glob(os.path.join(label_dir, "*.txt")):
with open(file, 'r') as f:
for line in f:
cls_id = int(line.strip().split()[0])
actual_classes.add(cls_id)
# 与配置文件对比
with open(data_yaml, 'r') as f:
config = yaml.safe_load(f)
if len(actual_classes) != config.get('nc', 0):
print(f"⚠️ 类别数不匹配: 实际={len(actual_classes)}, 配置={config.get('nc')}")
坑7:图片格式问题
问题: 某些图片无法加载
解决方案:
def validate_images(image_dir):
for img_file in glob.glob(os.path.join(image_dir, "*.*")):
try:
img = Image.open(img_file)
img.verify()
except Exception as e:
print(f"❌ 无效图片: {img_file}")
# 删除或标记无效图片
坑8:多GPU训练问题
问题: 使用多GPU时出现同步错误
解决方案:
# 使用单GPU或指定GPU
device = "0" # 使用第一个GPU
# 或使用CPU
device = "cpu"
坑9:数据增强过拟合
问题: 训练准确率高但验证准确率低
解决方案:
减少数据增强强度
增加验证集规模
使用早停(Early Stopping)
增加权重衰减
坑10:环境变量冲突
问题: 多个Python环境导致库冲突
解决方案:
# 使用conda环境隔离
conda create -n yolov8_fire python=3.10
conda activate yolov8_fire
pip install -r requirements.txt
# 或者在代码中指定Python路径
import sys
sys.path.insert(0, 'your_env_path/Lib/site-packages')
8. 性能优化:让训练更快更稳定
8.1 训练速度优化
1. 使用GPU加速
device = "0" if torch.cuda.is_available() else "cpu"
# 速度提升:GPU vs CPU = 10x ~ 50x
2. 调整batch size
# 根据显存调整
batch_size = 16 # 6GB显存
batch_size = 32 # 12GB显存
batch_size = 64 # 24GB显存
3. 使用混合精度训练
# 开启AMP(如果版本支持)
cmd += " amp=True"
# 速度提升:~30%,显存节省:~30%
4. 多线程数据加载
workers = 4 # 根据CPU核心数调整
# 减少数据加载等待时间
8.2 模型精度优化
1. 调整学习率
lr0 = 0.001 # 初始学习率
lrf = 0.01 # 最终学习率衰减因子
2. 使用预训练权重
pretrained = True # 使用COCO预训练权重
# 收敛更快,精度更高
3. 数据增强策略
# Mosaic增强(开启/关闭)
mosaic = 1.0 # 开启
mosaic = 0.0 # 关闭
# 翻转增强
fliplr = 0.5 # 水平翻转概率
flipud = 0.0 # 垂直翻转概率
8.3 显存优化技巧
1. 梯度累积
# 模拟更大的batch size
accumulate = 4 # 累积4次梯度更新
# 有效batch size = batch_size * accumulate
2. 使用梯度检查点
# 用时间换空间
torch.utils.checkpoint
3. 释放缓存
torch.cuda.empty_cache()
9. 完整代码解析:逐行带你读懂
9.1 项目结构
YOLOv8-Fire-Detection-GUI/
├── main.py # 主程序
├── datasets/ # 数据集
│ └── fire-8/
│ ├── train/
│ ├── valid/
│ ├── test/
│ └── data.yaml
├── runs/ # 训练输出
│ └── detect/
│ ├── train/
│ └── predict/
├── requirements.txt # 依赖列表
└── README.md # 说明文档
9.2 核心类解析
YOLOFireDetectionGUI
这是主窗口类,负责:
-
界面初始化和布局
-
用户交互事件处理
-
管理训练和推理流程
关键方法:
# 自动配置
def auto_configure(self):
"""启动时自动检测数据集和模型"""
# 1. 检测数据集
if os.path.exists(DATASET_PATH):
self.dataset_path.setText(DATASET_PATH)
self.console.append_success(f"Found dataset: {DATASET_PATH}")
# 2. 检测data.yaml
if os.path.exists(DATA_YAML):
self.data_yaml_path.setText(DATA_YAML)
self.console.append_success(f"Found data.yaml: {DATA_YAML}")
# 3. 检测训练好的模型
if os.path.exists(BEST_MODEL):
self.model_weights.setText(BEST_MODEL)
self.console.append_success(f"Found best model: {BEST_MODEL}")
# 4. 显示数据集统计
self.show_dataset_info()
# 开始训练
def start_training(self):
"""构建并执行训练命令"""
# 1. 验证输入
if not self.validate_inputs():
return
# 2. 构建命令
cmd = self.build_training_command()
# 3. 启动训练线程
self.training_worker = TrainingWorker(cmd, os.getcwd())
self.training_worker.output_signal.connect(self.console.append)
self.training_worker.progress_signal.connect(self.progress_bar.setValue)
self.training_worker.finished_signal.connect(self.training_finished)
self.training_worker.start()
TrainingWorker
这是训练工作线程,负责:
-
执行YOLO训练命令
-
实时捕获输出
-
解析训练进度
-
错误诊断
关键方法:
def run(self):
"""执行训练命令"""
# 1. 环境检查
self.check_environment()
# 2. 数据集检查
self.check_dataset()
# 3. 执行训练
self.process = subprocess.Popen(
self.command,
shell=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True
)
# 4. 实时读取输出
self.read_outputs()
# 5. 分析结果
self.analyze_result()
def read_outputs(self):
"""读取stdout和stderr"""
# 使用线程分别读取
stdout_thread = threading.Thread(target=self.read_stdout)
stderr_thread = threading.Thread(target=self.read_stderr)
stdout_thread.start()
stderr_thread.start()
def analyze_result(self):
"""智能分析训练结果"""
if return_code == 0:
self.handle_success()
else:
self.handle_failure()
9.3 样式系统
双主题支持:
# 暗黑主题
DARK_THEME = """
QMainWindow {
background-color: #1a1a2e;
}
QPushButton {
background-color: #0f3460;
color: white;
border-radius: 8px;
}
QPushButton:hover {
background-color: #16213e;
border: 1px solid #e94560;
}
"""
# 明亮主题
LIGHT_THEME = """
QMainWindow {
background-color: #f5f5f5;
}
QPushButton {
background-color: #e94560;
color: white;
border-radius: 8px;
}
QPushButton:hover {
background-color: #c73e54;
}
"""
9.4 错误处理
try:
# 执行可能出错的操作
result = subprocess.run(command, capture_output=True)
except subprocess.CalledProcessError as e:
# 处理子进程错误
self.console.append_error(f"子进程错误: {e}")
except FileNotFoundError as e:
# 处理文件不存在
self.console.append_error(f"文件不存在: {e}")
except Exception as e:
# 处理其他异常
self.console.append_error(f"未知错误: {e}")
self.console.append(traceback.format_exc())
10. 实战演示:从零训练一个火灾检测模型
10.1 准备数据集
步骤1:下载数据集
我使用Roboflow的火灾检测数据集:
from roboflow import Roboflow
rf = Roboflow(api_key="your_api_key")
project = rf.workspace("your_workspace").project("fire-wrpgm")
dataset = project.version(8).download("yolov8")
步骤2:检查数据集结构
import os
import glob
dataset_path = "datasets/fire-8"
for split in ['train', 'valid', 'test']:
images = glob.glob(os.path.join(dataset_path, split, 'images', '*'))
labels = glob.glob(os.path.join(dataset_path, split, 'labels', '*'))
print(f"{split}: {len(images)} images, {len(labels)} labels")
10.2 配置训练参数
打开GUI:
python main.py
参数配置:
-
Model Size: yolov8s.pt (small)
-
Epochs: 50
-
Image Size: 640
-
Batch Size: 16
-
Learning Rate: 0.001
-
Device: 0 (GPU)
10.3 开始训练
点击 "🚀 Start Training",训练开始!
训练过程输出:
[10:00:00] 🔍 开始环境检查…
[10:00:00] ✅ Ultralytics 版本: 8.0.20
[10:00:00] ✅ PyTorch 版本: 2.0.1
[10:00:00] ✅ CUDA 可用: NVIDIA GeForce RTX 3060
[10:00:00] 📂 检查数据集…
[10:00:00] ✅ data.yaml 存在
[10:00:00] ├─ train/: images=825, labels=825
[10:00:00] ├─ valid/: images=47, labels=47
[10:00:00] └─ test/: images=54, labels=54
[10:00:01] 🚀 开始执行训练…
[10:00:01] Model summary: 130 layers, 11,136,761 parameters
[10:00:02] Transferred 349/355 items from pretrained weights
[10:00:02] Epoch 1/50: 100% ━━━━━━━━━━━━ 52/52 [00:25<00:00]
[10:00:30] Epoch 1/50: 100% ━━━━━━━━━━━━ 52/52 [00:25<00:00, 2.08it/s]
10.4 查看训练结果
训练完成后,自动跳转到可视化页面,查看:
-
训练曲线(loss下降、精度上升)
-
混淆矩阵
-
验证集预测结果

10.5 模型推理
切换到 "🔍 Inference" 标签:
模型自动加载 best.pt
选择测试图片
点击 "🔍 Run Inference"
检测结果:
检测到: Fire (置信度: 0.92)
检测到: Smoke (置信度: 0.87)
10.6 性能评估
训练结果:
-
mAP50: 0.89
-
mAP50-95: 0.62
-
Precision: 0.91
-
Recall: 0.87
推理速度:
-
GPU: 30-50 FPS (RTX 3060)
-
CPU: 5-10 FPS (i7-10700)

11. 常见问题FAQ
Q1: 训练时GPU内存不足怎么办?
A: 按顺序尝试以下方法:
减小batch size(16 → 8 → 4)
减小image size(640 → 416)
关闭AMP(如果开启)
使用CPU训练(device=cpu)
Q2: 如何提高模型准确率?
A:
增加训练数据(至少1000张/类别)
增加训练轮数(50 → 100)
使用更大的模型(yolov8m.pt或yolov8l.pt)
调整学习率(0.001 → 0.0001)
数据增强优化
Q3: 推理速度太慢怎么办?
A:
使用nano模型(yolov8n.pt)
减小image size(640 → 320)
使用GPU推理
使用TensorRT加速
批量推理
Q4: 如何训练自己的数据集?
A:
按YOLOv8格式组织数据
创建data.yaml配置文件
在GUI中选择数据集路径
配置训练参数
开始训练
Q5: 程序崩溃了怎么办?
A:
查看 training_error_log.txt 文件
点击 "🐛 Debug Info" 查看系统状态
检查Python环境是否正确
重新安装依赖
在GitHub提Issue
Q6: 如何迁移到其他项目?
A:
修改数据集路径
修改类别名称
调整超参数
训练新模型
Q7: 支持哪些操作系统?
A:
-
Windows 10/11 ✅
-
Ubuntu 20.04+ ✅
-
macOS 10.15+ ✅
Q8: 需要联网吗?
A:
-
首次运行需要下载模型文件(约22MB)
-
之后可以离线使用
如果这篇文章对你有帮助,欢迎评论、点赞、收藏、转发! 🙏
你的支持是我继续创作的动力! 🚀
附A:
"""
YOLOv8 Early Fire Detection GUI – 增强修复版
添加了自动路径修复和错误恢复功能
"""
import sys
import os
import subprocess
import glob
import traceback
import yaml
import shutil
from datetime import datetime
from pathlib import Path
from PyQt5.QtWidgets import *
from PyQt5.QtCore import *
from PyQt5.QtGui import *
import torch
from PIL import Image
import matplotlib.pyplot as plt
from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas
# 项目路径配置
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
DATASET_PATH = os.path.join(PROJECT_ROOT, "datasets", "fire-8")
DATA_YAML = os.path.join(DATASET_PATH, "data.yaml")
BEST_MODEL = os.path.join(PROJECT_ROOT, "runs", "detect", "train", "weights", "best.pt")
TRAIN_RESULTS = os.path.join(PROJECT_ROOT, "runs", "detect", "train")
# ============ 数据集路径修复工具 ============
class DatasetPathFixer:
"""自动修复数据集路径问题"""
def __init__(self, dataset_path):
self.dataset_path = Path(dataset_path)
self.fixed = False
self.messages = []
def fix(self):
"""修复数据集路径"""
try:
# 查找 data.yaml
data_yaml = self._find_data_yaml()
if not data_yaml:
return False, "未找到 data.yaml 文件"
self.messages.append(f"找到 data.yaml: {data_yaml}")
# 读取配置
with open(data_yaml, 'r', encoding='utf-8') as f:
config = yaml.safe_load(f)
# 检查并修复路径
modified = False
for key in ['train', 'val', 'test']:
if key in config:
path = config[key]
if isinstance(path, str):
fixed_path = self._fix_path(path, key)
if fixed_path != path:
config[key] = fixed_path
modified = True
self.messages.append(f"修复 {key} 路径: {path} -> {fixed_path}")
if modified:
# 备份原文件
backup = data_yaml.with_suffix('.yaml.bak')
shutil.copy(data_yaml, backup)
self.messages.append(f"备份原文件: {backup}")
# 保存修改
with open(data_yaml, 'w', encoding='utf-8') as f:
yaml.dump(config, f, default_flow_style=False, allow_unicode=True)
self.fixed = True
return True, "数据集路径已修复"
return True, "数据集路径无需修复"
except Exception as e:
return False, f"修复失败: {str(e)}"
def _find_data_yaml(self):
"""查找 data.yaml 文件"""
# 检查根目录
data_yaml = self.dataset_path / 'data.yaml'
if data_yaml.exists():
return data_yaml
# 检查子目录
for subdir in self.dataset_path.iterdir():
if subdir.is_dir():
data_yaml = subdir / 'data.yaml'
if data_yaml.exists():
self.dataset_path = subdir
return data_yaml
# 递归搜索
for data_yaml in self.dataset_path.rglob('data.yaml'):
return data_yaml
return None
def _fix_path(self, path, split):
"""修复单个路径"""
# 如果是绝对路径且存在,直接返回
if os.path.isabs(path):
abs_path = Path(path)
if abs_path.exists():
return path
# 尝试相对路径
base_path = self.dataset_path
# 常见路径模式
patterns = [
path,
f'{split}/images',
f'../{split}/images',
f'{self.dataset_path.name}/{split}/images',
f'../{self.dataset_path.name}/{split}/images',
]
# 如果路径包含重复目录名,尝试修复
path_parts = Path(path).parts
if len(path_parts) >= 2 and path_parts[0] == path_parts[-1]:
patterns.append(f'{path_parts[-1]}/images')
patterns.append(f'{path_parts[-1]}')
for pattern in patterns:
full_path = base_path / pattern
if full_path.exists():
return str(full_path.absolute())
# 尝试在目录中搜索
for subdir in base_path.rglob('*'):
if subdir.is_dir() and subdir.name == split:
images_dir = subdir / 'images'
if images_dir.exists():
return str(images_dir.absolute())
if any(subdir.glob('*.[jJ][pP][gG]')):
return str(subdir.absolute())
# 递归搜索所有images目录
for images_dir in base_path.rglob('images'):
if images_dir.exists():
return str(images_dir.absolute())
return path
# ============ 样式表 ============
DARK_THEME = """
QMainWindow {
background-color: #1a1a2e;
}
QWidget {
background-color: #1a1a2e;
color: #e0e0e0;
font-family: 'Segoe UI', Arial, sans-serif;
}
QTabWidget::pane {
border: 2px solid #2d2d4a;
border-radius: 10px;
background-color: #16213e;
}
QTabBar::tab {
background-color: #1a1a2e;
color: #a0a0c0;
padding: 10px 20px;
border-top-left-radius: 8px;
border-top-right-radius: 8px;
margin-right: 2px;
font-weight: bold;
}
QTabBar::tab:selected {
background-color: #0f3460;
color: #ffffff;
border-bottom: 3px solid #e94560;
}
QPushButton {
background-color: #0f3460;
color: white;
border: none;
padding: 10px 20px;
border-radius: 8px;
font-weight: bold;
font-size: 13px;
}
QPushButton:hover {
background-color: #16213e;
border: 1px solid #e94560;
}
QPushButton:pressed {
background-color: #e94560;
}
QPushButton:disabled {
background-color: #2d2d4a;
color: #666;
}
QLineEdit, QTextEdit, QPlainTextEdit {
background-color: #16213e;
border: 2px solid #2d2d4a;
border-radius: 8px;
padding: 8px;
color: #e0e0e0;
}
QLineEdit:focus, QTextEdit:focus, QPlainTextEdit:focus {
border-color: #e94560;
}
QLabel {
color: #e0e0e0;
}
QGroupBox {
border: 2px solid #2d2d4a;
border-radius: 10px;
margin-top: 10px;
padding-top: 10px;
font-weight: bold;
}
QGroupBox::title {
subcontrol-origin: margin;
left: 10px;
padding: 0 10px 0 10px;
color: #e94560;
}
QProgressBar {
border: 2px solid #2d2d4a;
border-radius: 8px;
text-align: center;
background-color: #16213e;
}
QProgressBar::chunk {
background-color: #e94560;
border-radius: 6px;
}
QComboBox {
background-color: #16213e;
border: 2px solid #2d2d4a;
border-radius: 8px;
padding: 8px;
color: #e0e0e0;
}
QComboBox::drop-down {
border: none;
}
QComboBox::down-arrow {
image: none;
border-left: 5px solid transparent;
border-right: 5px solid transparent;
border-top: 5px solid #e0e0e0;
margin-right: 5px;
}
QComboBox QAbstractItemView {
background-color: #16213e;
border: 2px solid #2d2d4a;
selection-background-color: #0f3460;
}
QScrollBar:vertical {
border: none;
background: #1a1a2e;
width: 10px;
border-radius: 5px;
}
QScrollBar::handle:vertical {
background: #0f3460;
border-radius: 5px;
}
QScrollBar::handle:vertical:hover {
background: #e94560;
}
QScrollBar:horizontal {
border: none;
background: #1a1a2e;
height: 10px;
border-radius: 5px;
}
QScrollBar::handle:horizontal {
background: #0f3460;
border-radius: 5px;
}
QScrollBar::handle:horizontal:hover {
background: #e94560;
}
QMenuBar {
background-color: #1a1a2e;
color: #e0e0e0;
}
QMenuBar::item:selected {
background-color: #0f3460;
}
QMenu {
background-color: #1a1a2e;
border: 2px solid #2d2d4a;
}
QMenu::item:selected {
background-color: #0f3460;
}
QStatusBar {
background-color: #16213e;
color: #a0a0c0;
}
QCheckBox {
color: #e0e0e0;
}
QCheckBox::indicator {
width: 18px;
height: 18px;
border-radius: 4px;
border: 2px solid #2d2d4a;
background-color: #16213e;
}
QCheckBox::indicator:checked {
background-color: #e94560;
border-color: #e94560;
}
QSpinBox, QDoubleSpinBox {
background-color: #16213e;
border: 2px solid #2d2d4a;
border-radius: 8px;
padding: 5px;
color: #e0e0e0;
}
QSpinBox::up-button, QDoubleSpinBox::up-button,
QSpinBox::down-button, QDoubleSpinBox::down-button {
background-color: #0f3460;
border: none;
border-radius: 4px;
}
QSplitter::handle {
background-color: #2d2d4a;
}
"""
LIGHT_THEME = """
QMainWindow {
background-color: #f5f5f5;
}
QWidget {
background-color: #f5f5f5;
color: #333333;
font-family: 'Segoe UI', Arial, sans-serif;
}
QTabWidget::pane {
border: 2px solid #d0d0d0;
border-radius: 10px;
background-color: #ffffff;
}
QTabBar::tab {
background-color: #e8e8e8;
color: #555555;
padding: 10px 20px;
border-top-left-radius: 8px;
border-top-right-radius: 8px;
margin-right: 2px;
font-weight: bold;
}
QTabBar::tab:selected {
background-color: #ffffff;
color: #e94560;
border-bottom: 3px solid #e94560;
}
QPushButton {
background-color: #e94560;
color: white;
border: none;
padding: 10px 20px;
border-radius: 8px;
font-weight: bold;
font-size: 13px;
}
QPushButton:hover {
background-color: #c73e54;
}
QPushButton:pressed {
background-color: #a83245;
}
QPushButton:disabled {
background-color: #cccccc;
color: #888888;
}
QLineEdit, QTextEdit, QPlainTextEdit {
background-color: #ffffff;
border: 2px solid #d0d0d0;
border-radius: 8px;
padding: 8px;
color: #333333;
}
QLineEdit:focus, QTextEdit:focus, QPlainTextEdit:focus {
border-color: #e94560;
}
QLabel {
color: #333333;
}
QGroupBox {
border: 2px solid #d0d0d0;
border-radius: 10px;
margin-top: 10px;
padding-top: 10px;
font-weight: bold;
}
QGroupBox::title {
subcontrol-origin: margin;
left: 10px;
padding: 0 10px 0 10px;
color: #e94560;
}
QProgressBar {
border: 2px solid #d0d0d0;
border-radius: 8px;
text-align: center;
background-color: #ffffff;
}
QProgressBar::chunk {
background-color: #e94560;
border-radius: 6px;
}
QComboBox {
background-color: #ffffff;
border: 2px solid #d0d0d0;
border-radius: 8px;
padding: 8px;
color: #333333;
}
QComboBox::drop-down {
border: none;
}
QComboBox::down-arrow {
image: none;
border-left: 5px solid transparent;
border-right: 5px solid transparent;
border-top: 5px solid #333333;
margin-right: 5px;
}
QComboBox QAbstractItemView {
background-color: #ffffff;
border: 2px solid #d0d0d0;
selection-background-color: #e94560;
selection-color: white;
}
QScrollBar:vertical {
border: none;
background: #f0f0f0;
width: 10px;
border-radius: 5px;
}
QScrollBar::handle:vertical {
background: #d0d0d0;
border-radius: 5px;
}
QScrollBar::handle:vertical:hover {
background: #e94560;
}
QScrollBar:horizontal {
border: none;
background: #f0f0f0;
height: 10px;
border-radius: 5px;
}
QScrollBar::handle:horizontal {
background: #d0d0d0;
border-radius: 5px;
}
QScrollBar::handle:horizontal:hover {
background: #e94560;
}
QMenuBar {
background-color: #f5f5f5;
color: #333333;
}
QMenuBar::item:selected {
background-color: #e94560;
color: white;
}
QMenu {
background-color: #ffffff;
border: 2px solid #d0d0d0;
}
QMenu::item:selected {
background-color: #e94560;
color: white;
}
QStatusBar {
background-color: #ffffff;
color: #666666;
}
QCheckBox {
color: #333333;
}
QCheckBox::indicator {
width: 18px;
height: 18px;
border-radius: 4px;
border: 2px solid #d0d0d0;
background-color: #ffffff;
}
QCheckBox::indicator:checked {
background-color: #e94560;
border-color: #e94560;
}
QSpinBox, QDoubleSpinBox {
background-color: #ffffff;
border: 2px solid #d0d0d0;
border-radius: 8px;
padding: 5px;
color: #333333;
}
QSpinBox::up-button, QDoubleSpinBox::up-button,
QSpinBox::down-button, QDoubleSpinBox::down-button {
background-color: #e8e8e8;
border: none;
border-radius: 4px;
}
QSplitter::handle {
background-color: #d0d0d0;
}
"""
class ConsoleOutput(QPlainTextEdit):
"""Custom console widget for displaying command output"""
def __init__(self, parent=None):
super().__init__(parent)
self.setReadOnly(True)
self.setMaximumBlockCount(10000)
self.setFont(QFont("Consolas", 10))
def append(self, text):
"""Append text to console with timestamp"""
timestamp = datetime.now().strftime("%H:%M:%S")
self.appendPlainText(f"[{timestamp}] {text}")
self.verticalScrollBar().setValue(self.verticalScrollBar().maximum())
def append_error(self, text):
self.append(f"❌ {text}")
def append_warning(self, text):
self.append(f"⚠️ {text}")
def append_success(self, text):
self.append(f"✅ {text}")
def append_info(self, text):
self.append(f"ℹ️ {text}")
class TrainingWorker(QThread):
"""增强的Worker线程,带详细调试信息和错误恢复"""
output_signal = pyqtSignal(str)
finished_signal = pyqtSignal(int, str)
progress_signal = pyqtSignal(int)
def __init__(self, command, work_dir):
super().__init__()
self.command = command
self.work_dir = work_dir
self.process = None
self.is_running = True
self.error_output = []
self.full_output = []
def run(self):
try:
# ========== 1. 环境检查 ==========
self.output_signal.emit("=" * 60)
self.output_signal.emit("🔍 开始环境检查…")
self.output_signal.emit("=" * 60)
# 检查磁盘空间
try:
import psutil
disk = psutil.disk_usage(self.work_dir)
free_gb = disk.free / (1024**3)
self.output_signal.emit(f"💾 磁盘剩余空间: {free_gb:.1f} GB")
if free_gb < 2:
self.output_signal.emit("⚠️ 警告: 磁盘空间不足,建议至少有 2GB 空间")
except:
pass
# 检查 ultralytics
try:
import ultralytics
self.output_signal.emit(f"✅ Ultralytics 版本: {ultralytics.__version__}")
except ImportError as e:
self.output_signal.emit(f"❌ Ultralytics 未安装: {str(e)}")
self.finished_signal.emit(1, "Ultralytics not installed")
return
# 检查 PyTorch
try:
import torch
self.output_signal.emit(f"✅ PyTorch 版本: {torch.__version__}")
if torch.cuda.is_available():
self.output_signal.emit(f" ├─ CUDA 可用: {torch.cuda.get_device_name(0)}")
else:
self.output_signal.emit(" └─ ⚠️ CUDA 不可用, 使用 CPU")
except ImportError:
self.output_signal.emit("❌ PyTorch 未安装")
# ========== 2. 数据集检查 ==========
self.output_signal.emit("=" * 60)
self.output_signal.emit("📂 检查数据集…")
self.output_signal.emit("=" * 60)
# 解析 data.yaml 路径
data_yaml_path = None
if 'data=' in self.command:
parts = self.command.split('data=')
if len(parts) > 1:
data_yaml_path = parts[1].split()[0]
if data_yaml_path:
self.output_signal.emit(f"📄 data.yaml 路径: {data_yaml_path}")
if os.path.exists(data_yaml_path):
self.output_signal.emit(" ✅ data.yaml 存在")
# 尝试自动修复
try:
fixer = DatasetPathFixer(os.path.dirname(data_yaml_path))
success, msg = fixer.fix()
if success:
self.output_signal.emit(f" 🔧 {msg}")
else:
self.output_signal.emit(f" ⚠️ {msg}")
except Exception as e:
self.output_signal.emit(f" ⚠️ 修复尝试失败: {e}")
else:
self.output_signal.emit(f" ❌ data.yaml 不存在!")
self.finished_signal.emit(1, "data.yaml not found")
return
# ========== 3. 执行训练命令 ==========
self.output_signal.emit("=" * 60)
self.output_signal.emit("🚀 开始执行训练…")
self.output_signal.emit("=" * 60)
self.output_signal.emit(f"命令: {self.command}")
self.output_signal.emit("-" * 60)
# 设置环境变量
env = os.environ.copy()
env['PYTHONUNBUFFERED'] = '1'
env['TF_CPP_MIN_LOG_LEVEL'] = '2' # 减少 TensorFlow 日志
self.process = subprocess.Popen(
self.command,
shell=True,
cwd=self.work_dir,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
bufsize=1,
encoding='utf-8',
errors='replace',
env=env
)
self.output_signal.emit(f"✅ 进程已启动 (PID: {self.process.pid})")
# 使用线程读取输出
import threading
def read_stdout():
for line in self.process.stdout:
if not self.is_running:
break
line = line.strip()
if line:
self.full_output.append(line)
self.output_signal.emit(line)
# 解析进度
if 'epoch' in line.lower():
try:
import re
match = re.search(r'epoch[:\\s]+(\\d+)/(\\d+)', line, re.IGNORECASE)
if match:
current = int(match.group(1))
total = int(match.group(2))
if total > 0:
progress = int((current / total) * 100)
self.progress_signal.emit(progress)
except:
pass
# 检测错误关键字
error_keywords = ['error', 'failed', 'exception', 'traceback', 'cannot', 'unable']
if any(kw in line.lower() for kw in error_keywords):
self.error_output.append(f"ERROR: {line}")
def read_stderr():
for line in self.process.stderr:
if not self.is_running:
break
line = line.strip()
if line:
self.full_output.append(f"STDERR: {line}")
self.output_signal.emit(f"⚠️ STDERR: {line}")
self.error_output.append(line)
# 启动读取线程
stdout_thread = threading.Thread(target=read_stdout, daemon=True)
stderr_thread = threading.Thread(target=read_stderr, daemon=True)
stdout_thread.start()
stderr_thread.start()
# 等待进程完成
return_code = self.process.wait()
# 等待线程完成
stdout_thread.join(timeout=2)
stderr_thread.join(timeout=2)
# ========== 4. 结果分析 ==========
self.output_signal.emit("=" * 60)
self.output_signal.emit("📊 训练结果分析")
self.output_signal.emit("=" * 60)
if return_code == 0:
self.output_signal.emit("✅ 训练成功完成!")
self.finished_signal.emit(return_code, "Training completed successfully!")
else:
self.output_signal.emit(f"❌ 训练失败,返回码: {return_code}")
# 显示错误摘要
if self.error_output:
self.output_signal.emit("\\n📋 错误摘要 (最后10条):")
for i, err in enumerate(self.error_output[-10:], 1):
self.output_signal.emit(f" {i}. {err}")
# 分析常见错误
self.output_signal.emit("\\n🔍 错误分析:")
error_text = ' '.join(self.error_output).lower()
if 'cuda' in error_text or 'cudnn' in error_text:
self.output_signal.emit(" 💡 检测到 CUDA 相关错误,尝试:")
self.output_signal.emit(" 1. 设置 device=cpu")
self.output_signal.emit(" 2. 检查 CUDA 版本兼容性")
if 'out of memory' in error_text or 'oom' in error_text:
self.output_signal.emit(" 💡 检测到内存不足错误,尝试:")
self.output_signal.emit(" 1. 减少 batch size")
self.output_signal.emit(" 2. 减少 image size")
if 'file not found' in error_text or 'no such file' in error_text:
self.output_signal.emit(" 💡 检测到文件不存在错误,检查:")
self.output_signal.emit(" 1. data.yaml 中的路径是否正确")
self.output_signal.emit(" 2. 图片和标签文件是否存在")
if 'i/o operation on closed file' in error_text or 'closed file' in error_text:
self.output_signal.emit(" 💡 检测到文件 I/O 错误,尝试:")
self.output_signal.emit(" 1. 检查磁盘空间是否充足")
self.output_signal.emit(" 2. 以管理员身份运行程序")
self.output_signal.emit(" 3. 关闭杀毒软件或添加排除项")
self.output_signal.emit(" 4. 检查文件权限")
# 尝试修复
self.output_signal.emit("\\n🔧 尝试自动修复…")
try:
self._fix_file_io_error()
except Exception as e:
self.output_signal.emit(f"修复失败: {e}")
# 保存完整日志
log_file = os.path.join(self.work_dir, "training_error_log.txt")
try:
with open(log_file, 'w', encoding='utf-8') as f:
f.write("=" * 60 + "\\n")
f.write("训练错误日志\\n")
f.write(f"时间: {datetime.now()}\\n")
f.write(f"命令: {self.command}\\n")
f.write("=" * 60 + "\\n\\n")
f.write("完整输出:\\n")
f.write("\\n".join(self.full_output))
f.write("\\n\\n错误输出:\\n")
f.write("\\n".join(self.error_output))
self.output_signal.emit(f"\\n📄 完整错误日志已保存: {log_file}")
except Exception as e:
self.output_signal.emit(f"⚠️ 无法保存日志: {e}")
error_msg = f"Training failed with code {return_code}"
if self.error_output:
error_msg += f"\\nLast errors: {self.error_output[-3:]}"
self.finished_signal.emit(return_code, error_msg)
except FileNotFoundError as e:
self.output_signal.emit(f"❌ 命令未找到: {str(e)}")
self.output_signal.emit("💡 请确保 YOLO 命令行工具已安装并在 PATH 中")
self.finished_signal.emit(1, f"Command not found: {str(e)}")
except Exception as e:
self.output_signal.emit(f"❌ 未预期的错误: {str(e)}")
self.output_signal.emit(f"错误类型: {type(e).__name__}")
self.output_signal.emit(f"详细信息:\\n{traceback.format_exc()}")
self.finished_signal.emit(1, f"Unexpected error: {str(e)}")
def _fix_file_io_error(self):
"""修复文件 I/O 错误"""
# 检查并清理临时文件
runs_dir = Path(self.work_dir) / 'runs'
if runs_dir.exists():
# 删除临时文件
for temp_file in runs_dir.rglob('*.tmp'):
try:
temp_file.unlink()
self.output_signal.emit(f" 🗑️ 删除临时文件: {temp_file}")
except:
pass
# 检查权重目录权限
weights_dir = runs_dir / 'detect' / 'train' / 'weights'
if weights_dir.exists():
try:
# 尝试创建测试文件
test_file = weights_dir / '.test_write'
test_file.write_text('test')
test_file.unlink()
self.output_signal.emit(" ✅ 权重目录可写")
except:
self.output_signal.emit(" ⚠️ 权重目录不可写,请检查权限")
# 检查磁盘空间
try:
import psutil
disk = psutil.disk_usage(self.work_dir)
free_gb = disk.free / (1024**3)
self.output_signal.emit(f" 💾 当前磁盘剩余空间: {free_gb:.1f} GB")
if free_gb < 2:
self.output_signal.emit(" ⚠️ 磁盘空间不足!请清理磁盘或更换目录")
except:
pass
def stop(self):
self.is_running = False
if self.process:
self.output_signal.emit("⏹ 正在停止训练…")
self.process.terminate()
try:
self.process.wait(timeout=5)
self.output_signal.emit("✅ 训练已停止")
except:
self.process.kill()
self.output_signal.emit("⚠️ 强制终止训练")
class ImageViewer(QWidget):
"""Widget for displaying images with zoom and pan"""
def __init__(self, parent=None):
super().__init__(parent)
layout = QVBoxLayout(self)
layout.setContentsMargins(0, 0, 0, 0)
self.scene = QGraphicsScene()
self.view = QGraphicsView(self.scene)
self.view.setRenderHint(QPainter.Antialiasing)
self.view.setRenderHint(QPainter.SmoothPixmapTransform)
self.view.setDragMode(QGraphicsView.ScrollHandDrag)
self.view.setTransformationAnchor(QGraphicsView.AnchorUnderMouse)
self.view.setResizeAnchor(QGraphicsView.AnchorUnderMouse)
self.image_item = QGraphicsPixmapItem()
self.scene.addItem(self.image_item)
# Zoom controls
control_layout = QHBoxLayout()
zoom_in_btn = QPushButton("🔍+")
zoom_out_btn = QPushButton("🔍-")
fit_btn = QPushButton("Fit")
zoom_in_btn.clicked.connect(self.zoom_in)
zoom_out_btn.clicked.connect(self.zoom_out)
fit_btn.clicked.connect(self.fit_view)
control_layout.addWidget(zoom_in_btn)
control_layout.addWidget(zoom_out_btn)
control_layout.addWidget(fit_btn)
control_layout.addStretch()
layout.addWidget(self.view)
layout.addLayout(control_layout)
self.current_pixmap = None
self.zoom_factor = 1.0
def set_image(self, image_path):
if os.path.exists(image_path):
pixmap = QPixmap(image_path)
if not pixmap.isNull():
max_size = 800
if pixmap.width() > max_size or pixmap.height() > max_size:
pixmap = pixmap.scaled(
max_size, max_size,
Qt.KeepAspectRatio,
Qt.SmoothTransformation
)
self.current_pixmap = pixmap
self.image_item.setPixmap(pixmap)
self.fit_view()
return True
return False
def set_pixmap(self, pixmap):
self.current_pixmap = pixmap
self.image_item.setPixmap(pixmap)
self.fit_view()
def zoom_in(self):
self.zoom_factor *= 1.2
self.view.scale(1.2, 1.2)
def zoom_out(self):
self.zoom_factor *= 0.8
self.view.scale(0.8, 0.8)
def fit_view(self):
if self.image_item.pixmap() and not self.image_item.pixmap().isNull():
self.view.fitInView(self.image_item, Qt.KeepAspectRatio)
self.zoom_factor = 1.0
class MatplotlibWidget(FigureCanvas):
"""Widget for displaying matplotlib figures"""
def __init__(self, parent=None, width=5, height=4, dpi=100):
self.figure = plt.figure(figsize=(width, height), dpi=dpi)
self.figure.set_facecolor('#16213e' if parent and hasattr(parent, 'is_dark') else '#ffffff')
super().__init__(self.figure)
self.setParent(parent)
self.axes = self.figure.add_subplot(111)
self.axes.set_facecolor('#1a1a2e' if parent and hasattr(parent, 'is_dark') else '#f5f5f5')
self.axes.axis('off')
def plot_image(self, image_path):
if os.path.exists(image_path):
try:
img = Image.open(image_path)
self.axes.clear()
self.axes.imshow(img)
self.axes.axis('off')
self.draw()
return True
except Exception as e:
print(f"Error plotting: {e}")
return False
class ClickableCard(QPushButton):
"""Custom clickable card widget"""
def __init__(self, title, description, icon, parent=None):
super().__init__(parent)
self.setFixedHeight(150)
self.setStyleSheet("""
QPushButton {
background-color: #16213e;
border: 2px solid #2d2d4a;
border-radius: 15px;
text-align: left;
padding: 15px;
}
QPushButton:hover {
border-color: #e94560;
background-color: #1a1a2e;
}
""")
layout = QVBoxLayout(self)
layout.setSpacing(8)
icon_label = QLabel(icon)
icon_label.setStyleSheet("font-size: 36px;")
layout.addWidget(icon_label)
title_label = QLabel(title)
title_label.setStyleSheet("font-size: 18px; font-weight: bold; color: #e94560;")
layout.addWidget(title_label)
desc_label = QLabel(description)
desc_label.setStyleSheet("font-size: 12px; color: #a0a0c0;")
desc_label.setWordWrap(True)
layout.addWidget(desc_label)
class YOLOFireDetectionGUI(QMainWindow):
"""Main application window"""
def __init__(self):
super().__init__()
self.is_dark = True
self.training_worker = None
self.viz_image_path = None
self.init_ui()
self.apply_theme()
self.auto_configure()
def auto_configure(self):
"""自动检测并配置数据集和模型路径,同时修复路径问题"""
# 检查并修复数据集
if os.path.exists(DATASET_PATH):
self.dataset_path.setText(DATASET_PATH)
# 尝试修复数据集路径
try:
fixer = DatasetPathFixer(DATASET_PATH)
success, msg = fixer.fix()
if success:
self.console.append_success(f"数据集路径修复: {msg}")
else:
self.console.append_warning(f"数据集路径修复失败: {msg}")
except Exception as e:
self.console.append_warning(f"自动修复异常: {e}")
# 显示数据集信息
self.show_dataset_info()
else:
self.console.append_warning(f"数据集目录不存在: {DATASET_PATH}")
# 检查 data.yaml
if os.path.exists(DATA_YAML):
self.data_yaml_path.setText(DATA_YAML)
self.console.append_success(f"找到 data.yaml: {DATA_YAML}")
else:
self.console.append_warning(f"data.yaml 未找到: {DATA_YAML}")
# 检查训练好的模型
if os.path.exists(BEST_MODEL):
self.model_weights.setText(BEST_MODEL)
self.console.append_success(f"找到最佳模型: {BEST_MODEL}")
# 加载训练结果
results_path = os.path.join(TRAIN_RESULTS, "results.png")
if os.path.exists(results_path):
self.results_viewer.set_image(results_path)
self.console.append_success("已加载训练结果")
else:
self.console.append_warning("未找到训练好的模型,请先训练模型")
# 显示GPU信息
self.show_gpu_info()
def show_dataset_info(self):
"""显示数据集信息"""
if not os.path.exists(DATASET_PATH):
return
info = f"\\n📊 数据集信息:"
# 检查 data.yaml
if os.path.exists(DATA_YAML):
try:
with open(DATA_YAML, 'r', encoding='utf-8') as f:
config = yaml.safe_load(f)
info += f"\\n ├─ 类别数: {config.get('nc', '未知')}"
info += f"\\n ├─ 类别名称: {config.get('names', [])}"
except:
pass
# 统计图片
for split in ['train', 'valid', 'test']:
split_path = os.path.join(DATASET_PATH, split, 'images')
if os.path.exists(split_path):
count = len(glob.glob(os.path.join(split_path, "*.*")))
info += f"\\n ├─ {split}: {count} 张图片"
else:
# 尝试其他可能的路径
alt_path = os.path.join(DATASET_PATH, split)
if os.path.exists(alt_path):
count = len(glob.glob(os.path.join(alt_path, "*.[jJ][pP][gG]")))
info += f"\\n ├─ {split}: {count} 张图片"
else:
info += f"\\n ├─ {split}: 未找到"
self.console.append(info)
def show_gpu_info(self):
"""显示GPU信息"""
try:
import torch
if torch.cuda.is_available():
gpu_name = torch.cuda.get_device_name(0)
gpu_memory = torch.cuda.get_device_properties(0).total_memory / (1024**3)
self.console.append_success(f"GPU: {gpu_name} ({gpu_memory:.1f} GB)")
# 更新设备选择
self.device.setCurrentIndex(0)
else:
self.console.append_warning("GPU 不可用,将使用 CPU 训练")
self.device.setCurrentIndex(1)
except:
pass
def init_ui(self):
self.setWindowTitle("🔥 YOLOv8 Early Fire Detection – 增强版")
self.setGeometry(100, 100, 1400, 900)
self.setWindowIcon(QIcon())
central_widget = QWidget()
self.setCentralWidget(central_widget)
main_layout = QVBoxLayout(central_widget)
main_layout.setSpacing(10)
main_layout.setContentsMargins(15, 15, 15, 15)
self.create_menu_bar()
self.create_toolbar()
self.tabs = QTabWidget()
self.tabs.setDocumentMode(True)
self.setup_welcome_tab()
self.setup_training_tab()
self.setup_inference_tab()
self.setup_visualization_tab()
self.setup_console_tab()
main_layout.addWidget(self.tabs)
self.status_bar = QStatusBar()
self.setStatusBar(self.status_bar)
self.status_bar.showMessage("就绪")
self.progress_bar = QProgressBar()
self.progress_bar.setMaximumWidth(200)
self.progress_bar.setVisible(False)
self.status_bar.addPermanentWidget(self.progress_bar)
def create_menu_bar(self):
menubar = self.menuBar()
file_menu = menubar.addMenu("文件")
exit_action = QAction("退出", self)
exit_action.setShortcut("Ctrl+Q")
exit_action.triggered.connect(self.close)
file_menu.addAction(exit_action)
view_menu = menubar.addMenu("视图")
theme_action = QAction("切换主题", self)
theme_action.setShortcut("Ctrl+T")
theme_action.triggered.connect(self.toggle_theme)
view_menu.addAction(theme_action)
tools_menu = menubar.addMenu("工具")
fix_dataset_action = QAction("修复数据集路径", self)
fix_dataset_action.triggered.connect(self.fix_dataset_path)
tools_menu.addAction(fix_dataset_action)
check_disk_action = QAction("检查磁盘空间", self)
check_disk_action.triggered.connect(self.check_disk_space)
tools_menu.addAction(check_disk_action)
help_menu = menubar.addMenu("帮助")
about_action = QAction("关于", self)
about_action.triggered.connect(self.show_about)
help_menu.addAction(about_action)
def create_toolbar(self):
toolbar = self.addToolBar("主工具栏")
toolbar.setMovable(False)
theme_btn = QAction("🌓 主题", self)
theme_btn.triggered.connect(self.toggle_theme)
toolbar.addAction(theme_btn)
toolbar.addSeparator()
train_btn = QAction("🚀 训练", self)
train_btn.triggered.connect(lambda: self.tabs.setCurrentIndex(1))
toolbar.addAction(train_btn)
infer_btn = QAction("🔍 推理", self)
infer_btn.triggered.connect(lambda: self.tabs.setCurrentIndex(2))
toolbar.addAction(infer_btn)
toolbar.addSeparator()
clear_btn = QAction("🗑️ 清空控制台", self)
clear_btn.triggered.connect(self.clear_console)
toolbar.addAction(clear_btn)
debug_btn = QAction("🐛 调试信息", self)
debug_btn.triggered.connect(self.show_debug_info)
toolbar.addAction(debug_btn)
def fix_dataset_path(self):
"""修复数据集路径"""
path = self.dataset_path.text().strip()
if not path:
QMessageBox.warning(self, "警告", "请先选择数据集路径")
return
if not os.path.exists(path):
QMessageBox.warning(self, "警告", f"路径不存在: {path}")
return
self.console.append("🔧 开始修复数据集路径…")
fixer = DatasetPathFixer(path)
success, msg = fixer.fix()
if success:
self.console.append_success(msg)
# 刷新显示
self.show_dataset_info()
else:
self.console.append_error(msg)
QMessageBox.warning(self, "修复失败", msg)
def check_disk_space(self):
"""检查磁盘空间"""
try:
import psutil
disk = psutil.disk_usage(os.getcwd())
free_gb = disk.free / (1024**3)
total_gb = disk.total / (1024**3)
used_gb = disk.used / (1024**3)
msg = f"""📊 磁盘空间信息:
总容量: {total_gb:.1f} GB
已用空间: {used_gb:.1f} GB
剩余空间: {free_gb:.1f} GB
使用率: {disk.percent}%
{'⚠️ 警告: 磁盘空间不足,建议清理!' if free_gb < 2 else '✅ 磁盘空间充足'}
"""
self.console.append(msg)
QMessageBox.information(self, "磁盘空间", msg)
except Exception as e:
self.console.append_error(f"检查磁盘空间失败: {e}")
def show_debug_info(self):
"""显示调试信息"""
self.console.append("=" * 60)
self.console.append("🐛 调试信息")
self.console.append("=" * 60)
self.console.append(f"Python 版本: {sys.version}")
self.console.append(f"工作目录: {os.getcwd()}")
# 检查关键文件
self.console.append("\\n📁 关键文件检查:")
files_to_check = [
("data.yaml", DATA_YAML),
("数据集目录", DATASET_PATH),
("最佳模型", BEST_MODEL),
]
for name, path in files_to_check:
exists = "✅" if os.path.exists(path) else "❌"
self.console.append(f" {exists} {name}: {path}")
# 检查 ultralytics
try:
import ultralytics
self.console.append(f"\\n✅ ultralytics 版本: {ultralytics.__version__}")
except ImportError:
self.console.append("\\n❌ ultralytics 未安装")
self.console.append("\\n" + "=" * 60)
def setup_welcome_tab(self):
tab = QWidget()
layout = QVBoxLayout(tab)
layout.setSpacing(20)
layout.setContentsMargins(30, 30, 30, 30)
title = QLabel("🔥 YOLOv8 Early Fire Detection")
title.setStyleSheet("font-size: 32px; font-weight: bold; color: #e94560;")
title.setAlignment(Qt.AlignCenter)
layout.addWidget(title)
subtitle = QLabel("Train custom fire detection models with YOLOv8")
subtitle.setStyleSheet("font-size: 18px; color: #a0a0c0;")
subtitle.setAlignment(Qt.AlignCenter)
layout.addWidget(subtitle)
layout.addSpacing(20)
cards_layout = QHBoxLayout()
cards_layout.setSpacing(20)
info_cards = [
("📊 Train", "Train YOLOv8 on your custom fire detection dataset", "📊", 1),
("🔍 Infer", "Run inference with trained models on images/videos", "🔍", 2),
("📈 Visualize", "View training results and predictions", "📈", 3),
]
for title_text, desc, icon, tab_idx in info_cards:
card = ClickableCard(title_text, desc, icon)
card.clicked.connect(lambda checked, idx=tab_idx: self.tabs.setCurrentIndex(idx))
cards_layout.addWidget(card)
layout.addLayout(cards_layout)
layout.addSpacing(20)
guide_group = QGroupBox("快速开始指南")
guide_layout = QVBoxLayout(guide_group)
guide_text = QTextEdit()
guide_text.setReadOnly(True)
guide_text.setHtml("""
<h3>快速开始</h3>
<ol>
<li><b>准备数据集:</b> 确保数据集在 datasets/fire-8 目录下</li>
<li><b>配置训练:</b> 点击"训练"标签页,设置训练参数</li>
<li><b>训练模型:</b> 点击"开始训练",在控制台查看进度</li>
<li><b>推理测试:</b> 点击"推理"标签页,选择模型和输入源</li>
</ol>
<p style="margin-top: 10px;">
<b>💡 提示:</b> 如果遇到路径错误,可以使用"工具"菜单中的"修复数据集路径"功能
</p>
""")
guide_layout.addWidget(guide_text)
layout.addWidget(guide_group)
layout.addStretch()
self.tabs.addTab(tab, "🏠 欢迎")
def setup_training_tab(self):
tab = QWidget()
layout = QHBoxLayout(tab)
layout.setSpacing(15)
# Left panel
left_panel = QWidget()
left_panel.setMaximumWidth(550)
left_layout = QVBoxLayout(left_panel)
left_layout.setSpacing(15)
# Dataset configuration
dataset_group = QGroupBox("📁 数据集配置")
dataset_layout = QGridLayout(dataset_group)
dataset_layout.setSpacing(10)
dataset_layout.addWidget(QLabel("数据集目录:"), 0, 0)
self.dataset_path = QLineEdit()
self.dataset_path.setPlaceholderText("选择包含 train/val/test 的数据集目录")
dataset_layout.addWidget(self.dataset_path, 0, 1)
browse_dataset_btn = QPushButton("📂 浏览")
browse_dataset_btn.clicked.connect(self.browse_dataset)
dataset_layout.addWidget(browse_dataset_btn, 0, 2)
dataset_layout.addWidget(QLabel("data.yaml:"), 1, 0)
self.data_yaml_path = QLineEdit()
self.data_yaml_path.setPlaceholderText("选择 data.yaml 文件")
dataset_layout.addWidget(self.data_yaml_path, 1, 1)
browse_yaml_btn = QPushButton("📄 浏览")
browse_yaml_btn.clicked.connect(self.browse_data_yaml)
dataset_layout.addWidget(browse_yaml_btn, 1, 2)
# 添加修复按钮
fix_path_btn = QPushButton("🔧 修复路径")
fix_path_btn.clicked.connect(self.fix_dataset_path)
dataset_layout.addWidget(fix_path_btn, 2, 0, 1, 3)
left_layout.addWidget(dataset_group)
# Model configuration
model_group = QGroupBox("🤖 模型配置")
model_layout = QGridLayout(model_group)
model_layout.setSpacing(10)
model_layout.addWidget(QLabel("模型大小:"), 0, 0)
self.model_size = QComboBox()
self.model_size.addItems(["yolov8n.pt (nano)", "yolov8s.pt (small)", "yolov8m.pt (medium)",
"yolov8l.pt (large)", "yolov8x.pt (x-large)"])
self.model_size.setCurrentIndex(1)
model_layout.addWidget(self.model_size, 0, 1)
model_layout.addWidget(QLabel("预训练:"), 1, 0)
self.pretrained_check = QCheckBox("使用预训练权重")
self.pretrained_check.setChecked(True)
model_layout.addWidget(self.pretrained_check, 1, 1)
left_layout.addWidget(model_group)
# Training parameters
params_group = QGroupBox("⚙️ 训练参数")
params_layout = QGridLayout(params_group)
params_layout.setSpacing(10)
params_layout.addWidget(QLabel("训练轮数:"), 0, 0)
self.epochs = QSpinBox()
self.epochs.setRange(1, 1000)
self.epochs.setValue(50)
params_layout.addWidget(self.epochs, 0, 1)
params_layout.addWidget(QLabel("图片尺寸:"), 1, 0)
self.img_size = QSpinBox()
self.img_size.setRange(320, 1280)
self.img_size.setSingleStep(32)
self.img_size.setValue(640)
params_layout.addWidget(self.img_size, 1, 1)
params_layout.addWidget(QLabel("批次大小:"), 2, 0)
self.batch_size = QSpinBox()
self.batch_size.setRange(1, 64)
self.batch_size.setValue(16)
params_layout.addWidget(self.batch_size, 2, 1)
params_layout.addWidget(QLabel("学习率:"), 3, 0)
self.lr = QDoubleSpinBox()
self.lr.setRange(0.0001, 0.1)
self.lr.setSingleStep(0.0001)
self.lr.setValue(0.001)
self.lr.setDecimals(4)
params_layout.addWidget(self.lr, 3, 1)
params_layout.addWidget(QLabel("工作进程:"), 4, 0)
self.workers = QSpinBox()
self.workers.setRange(0, 16)
self.workers.setValue(4)
params_layout.addWidget(self.workers, 4, 1)
params_layout.addWidget(QLabel("设备:"), 5, 0)
self.device = QComboBox()
self.device.addItems(["0 (GPU)", "cpu (CPU)"])
params_layout.addWidget(self.device, 5, 1)
left_layout.addWidget(params_group)
# Training controls
controls_group = QGroupBox("🎮 训练控制")
controls_layout = QVBoxLayout(controls_group)
controls_row = QHBoxLayout()
self.train_btn = QPushButton("🚀 开始训练")
self.train_btn.clicked.connect(self.start_training)
self.train_btn.setStyleSheet("""
QPushButton {
background-color: #e94560;
font-size: 14px;
padding: 12px 30px;
}
QPushButton:hover {
background-color: #c73e54;
}
""")
controls_row.addWidget(self.train_btn)
self.stop_btn = QPushButton("⏹ 停止训练")
self.stop_btn.clicked.connect(self.stop_training)
self.stop_btn.setEnabled(False)
self.stop_btn.setStyleSheet("""
QPushButton {
background-color: #ff6b6b;
font-size: 14px;
padding: 12px 30px;
}
QPushButton:hover {
background-color: #e55555;
}
""")
controls_row.addWidget(self.stop_btn)
controls_row.addStretch()
controls_layout.addLayout(controls_row)
options_row = QHBoxLayout()
self.validate_check = QCheckBox("训练后验证")
self.validate_check.setChecked(True)
options_row.addWidget(self.validate_check)
options_row.addStretch()
controls_layout.addLayout(options_row)
left_layout.addWidget(controls_group)
left_layout.addStretch()
# Right panel – Training results
right_panel = QWidget()
right_layout = QVBoxLayout(right_panel)
right_layout.setSpacing(10)
results_group = QGroupBox("📊 训练结果")
results_layout = QVBoxLayout(results_group)
self.results_viewer = ImageViewer()
results_layout.addWidget(self.results_viewer)
results_controls = QHBoxLayout()
self.results_combo = QComboBox()
self.results_combo.addItems([
"results.png", "confusion_matrix.png",
"val_batch0_pred.jpg", "val_batch0_labels.jpg"
])
self.results_combo.currentTextChanged.connect(self.load_training_result)
results_controls.addWidget(QLabel("查看:"))
results_controls.addWidget(self.results_combo)
results_controls.addStretch()
results_layout.addLayout(results_controls)
right_layout.addWidget(results_group)
layout.addWidget(left_panel)
layout.addWidget(right_panel)
self.tabs.addTab(tab, "🚀 训练")
def setup_inference_tab(self):
tab = QWidget()
layout = QHBoxLayout(tab)
layout.setSpacing(15)
# Left panel
left_panel = QWidget()
left_panel.setMaximumWidth(500)
left_layout = QVBoxLayout(left_panel)
left_layout.setSpacing(15)
# Model selection
model_group = QGroupBox("📂 模型选择")
model_layout = QGridLayout(model_group)
model_layout.setSpacing(10)
model_layout.addWidget(QLabel("模型权重 (.pt):"), 0, 0)
self.model_weights = QLineEdit()
self.model_weights.setPlaceholderText("选择训练好的模型权重文件")
model_layout.addWidget(self.model_weights, 0, 1)
browse_model_btn = QPushButton("📂 浏览")
browse_model_btn.clicked.connect(self.browse_model_weights)
model_layout.addWidget(browse_model_btn, 0, 2)
use_best_btn = QPushButton("使用最佳模型")
use_best_btn.clicked.connect(self.use_best_model)
model_layout.addWidget(use_best_btn, 1, 0, 1, 3)
left_layout.addWidget(model_group)
# Inference source
source_group = QGroupBox("📷 推理源")
source_layout = QGridLayout(source_group)
source_layout.setSpacing(10)
source_layout.addWidget(QLabel("源:"), 0, 0)
self.source_path = QLineEdit()
self.source_path.setPlaceholderText("图片路径、视频路径或目录")
source_layout.addWidget(self.source_path, 0, 1)
browse_source_btn = QPushButton("📂 浏览")
browse_source_btn.clicked.connect(self.browse_source)
source_layout.addWidget(browse_source_btn, 0, 2)
source_layout.addWidget(QLabel("置信度:"), 1, 0)
self.confidence = QDoubleSpinBox()
self.confidence.setRange(0.0, 1.0)
self.confidence.setSingleStep(0.05)
self.confidence.setValue(0.25)
source_layout.addWidget(self.confidence, 1, 1, 1, 2)
source_layout.addWidget(QLabel("源类型:"), 2, 0)
self.source_type = QComboBox()
self.source_type.addItems(["image", "video", "webcam", "directory"])
source_layout.addWidget(self.source_type, 2, 1, 1, 2)
left_layout.addWidget(source_group)
# Inference controls
controls_group = QGroupBox("🎮 推理控制")
controls_layout = QVBoxLayout(controls_group)
self.infer_btn = QPushButton("🔍 运行推理")
self.infer_btn.clicked.connect(self.run_inference)
self.infer_btn.setStyleSheet("""
QPushButton {
background-color: #4ecdc4;
font-size: 14px;
padding: 12px 30px;
}
QPushButton:hover {
background-color: #3dbdb4;
}
""")
controls_layout.addWidget(self.infer_btn)
options_row = QHBoxLayout()
self.save_results_check = QCheckBox("保存结果")
self.save_results_check.setChecked(True)
options_row.addWidget(self.save_results_check)
options_row.addStretch()
controls_layout.addLayout(options_row)
left_layout.addWidget(controls_group)
left_layout.addStretch()
# Right panel
right_panel = QWidget()
right_layout = QVBoxLayout(right_panel)
right_layout.setSpacing(10)
inference_group = QGroupBox("🔍 推理结果")
inference_layout = QVBoxLayout(inference_group)
self.inference_viewer = ImageViewer()
inference_layout.addWidget(self.inference_viewer)
self.inference_info = QLabel("准备推理")
self.inference_info.setAlignment(Qt.AlignCenter)
self.inference_info.setStyleSheet("color: #a0a0c0; font-size: 12px;")
inference_layout.addWidget(self.inference_info)
right_layout.addWidget(inference_group)
layout.addWidget(left_panel)
layout.addWidget(right_panel)
self.tabs.addTab(tab, "🔍 推理")
def setup_visualization_tab(self):
tab = QWidget()
layout = QVBoxLayout(tab)
layout.setSpacing(15)
controls_layout = QHBoxLayout()
controls_layout.addWidget(QLabel("查看:"))
self.viz_combo = QComboBox()
self.viz_combo.addItems([
"训练结果", "混淆矩阵",
"验证预测", "自定义图片"
])
self.viz_combo.currentTextChanged.connect(self.update_visualization)
controls_layout.addWidget(self.viz_combo)
self.viz_path_btn = QPushButton("📂 浏览图片")
self.viz_path_btn.clicked.connect(self.browse_viz_image)
controls_layout.addWidget(self.viz_path_btn)
controls_layout.addStretch()
layout.addLayout(controls_layout)
self.viz_widget = MatplotlibWidget(self, width=8, height=6, dpi=100)
layout.addWidget(self.viz_widget)
self.viz_info = QLabel("请选择可视化选项")
self.viz_info.setAlignment(Qt.AlignCenter)
self.viz_info.setStyleSheet("color: #a0a0c0; font-size: 12px; padding: 5px;")
layout.addWidget(self.viz_info)
self.tabs.addTab(tab, "📈 可视化")
def setup_console_tab(self):
tab = QWidget()
layout = QVBoxLayout(tab)
layout.setContentsMargins(0, 0, 0, 0)
toolbar = QHBoxLayout()
clear_btn = QPushButton("🗑️ 清空")
clear_btn.clicked.connect(self.clear_console)
toolbar.addWidget(clear_btn)
save_btn = QPushButton("💾 保存日志")
save_btn.clicked.connect(self.save_console_log)
toolbar.addWidget(save_btn)
toolbar.addStretch()
layout.addLayout(toolbar)
self.console = ConsoleOutput()
layout.addWidget(self.console)
self.tabs.addTab(tab, "📟 控制台")
def apply_theme(self):
if self.is_dark:
self.setStyleSheet(DARK_THEME)
else:
self.setStyleSheet(LIGHT_THEME)
def toggle_theme(self):
self.is_dark = not self.is_dark
self.apply_theme()
def show_about(self):
QMessageBox.about(
self,
"关于 YOLOv8 火灾检测",
"""
<h2>🔥 YOLOv8 早期火灾检测</h2>
<p>版本 2.0.0</p>
<p>训练和部署 YOLOv8 模型用于早期火灾检测</p>
<br>
<p><b>功能:</b></p>
<ul>
<li>自定义数据集训练</li>
<li>实时推理</li>
<li>可视化工具</li>
<li>GPU 加速支持</li>
<li>自动路径修复</li>
</ul>
<br>
<p>基于 PyQt5 和 Ultralytics YOLOv8</p>
"""
)
def browse_dataset(self):
path = QFileDialog.getExistingDirectory(
self,
"选择数据集目录",
DATASET_PATH if os.path.exists(DATASET_PATH) else "",
QFileDialog.ShowDirsOnly
)
if path:
self.dataset_path.setText(path)
yaml_path = os.path.join(path, "data.yaml")
if os.path.exists(yaml_path):
self.data_yaml_path.setText(yaml_path)
self.console.append_success(f"找到 data.yaml")
else:
self.console.append_warning(f"data.yaml 未找到,请手动选择")
def browse_data_yaml(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择 data.yaml",
DATA_YAML if os.path.exists(DATA_YAML) else "",
"YAML Files (*.yaml *.yml)"
)
if path:
self.data_yaml_path.setText(path)
def browse_model_weights(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择模型权重",
BEST_MODEL if os.path.exists(BEST_MODEL) else "",
"PyTorch Weights (*.pt)"
)
if path:
self.model_weights.setText(path)
def browse_source(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择源", "", "All Files (*.*)"
)
if path:
self.source_path.setText(path)
def browse_viz_image(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择图片", "", "Image Files (*.jpg *.jpeg *.png *.bmp)"
)
if path:
self.viz_image_path = path
self.viz_combo.setCurrentText("自定义图片")
self.update_visualization()
def use_best_model(self):
if os.path.exists(BEST_MODEL):
self.model_weights.setText(BEST_MODEL)
self.console.append_success("使用训练好的最佳模型")
else:
QMessageBox.warning(self, "警告", "最佳模型未找到,请先训练模型")
def clear_console(self):
self.console.clear()
def save_console_log(self):
path, _ = QFileDialog.getSaveFileName(
self, "保存控制台日志", "", "Text Files (*.txt)"
)
if path:
with open(path, 'w', encoding='utf-8') as f:
f.write(self.console.toPlainText())
QMessageBox.information(self, "成功", f"日志已保存到 {path}")
def load_training_result(self, filename):
if not filename:
return
image_path = os.path.join(TRAIN_RESULTS, filename)
if os.path.exists(image_path):
self.results_viewer.set_image(image_path)
else:
self.console.append_warning(f"文件未找到: {image_path}")
def update_visualization(self):
selection = self.viz_combo.currentText()
if selection == "训练结果":
path = os.path.join(TRAIN_RESULTS, "results.png")
if os.path.exists(path):
self.viz_widget.plot_image(path)
self.viz_info.setText("📊 训练结果概览")
else:
self.viz_info.setText("⚠️ 训练结果未找到,请先训练模型")
elif selection == "混淆矩阵":
path = os.path.join(TRAIN_RESULTS, "confusion_matrix.png")
if os.path.exists(path):
self.viz_widget.plot_image(path)
self.viz_info.setText("📊 混淆矩阵")
else:
self.viz_info.setText("⚠️ 混淆矩阵未找到")
elif selection == "验证预测":
path = os.path.join(TRAIN_RESULTS, "val_batch0_pred.jpg")
if os.path.exists(path):
self.viz_widget.plot_image(path)
self.viz_info.setText("🔍 验证批次预测")
else:
self.viz_info.setText("⚠️ 验证预测未找到")
elif selection == "自定义图片":
if hasattr(self, 'viz_image_path') and self.viz_image_path and os.path.exists(self.viz_image_path):
self.viz_widget.plot_image(self.viz_image_path)
self.viz_info.setText(f"🖼️ 自定义图片: {os.path.basename(self.viz_image_path)}")
else:
self.viz_info.setText("📷 使用浏览按钮选择自定义图片")
def start_training(self):
# 验证输入
if not self.data_yaml_path.text():
QMessageBox.warning(self, "警告", "请选择 data.yaml 文件")
return
if not os.path.exists(self.data_yaml_path.text()):
QMessageBox.warning(self, "警告", "data.yaml 文件不存在")
return
# 检查磁盘空间
try:
import psutil
disk = psutil.disk_usage(os.getcwd())
free_gb = disk.free / (1024**3)
if free_gb < 2:
reply = QMessageBox.warning(
self, "磁盘空间不足",
f"磁盘剩余空间仅 {free_gb:.1f} GB,建议至少有 2GB 空间。\\n是否继续?",
QMessageBox.Yes | QMessageBox.No
)
if reply == QMessageBox.No:
return
except:
pass
# 构建命令
model_text = self.model_size.currentText()
model = model_text.split()[0]
data = self.data_yaml_path.text()
epochs = self.epochs.value()
imgsz = self.img_size.value()
batch = self.batch_size.value()
lr = self.lr.value()
workers = self.workers.value()
device = self.device.currentText().split()[0]
# 使用更稳定的设置
cmd = (f"yolo task=detect mode=train model={model} data={data} "
f"epochs={epochs} imgsz={imgsz} batch={batch} lr0={lr} "
f"workers={workers} device={device} amp=False save=True exist_ok=True")
if self.pretrained_check.isChecked():
cmd += " pretrained=True"
cmd += " plots=True"
self.console.append(f"🚀 开始训练:\\n{cmd}\\n{'-'*50}")
# 禁用控件
self.train_btn.setEnabled(False)
self.stop_btn.setEnabled(True)
self.progress_bar.setVisible(True)
self.progress_bar.setValue(0)
self.status_bar.showMessage("训练中…")
# 启动训练线程
self.training_worker = TrainingWorker(cmd, os.getcwd())
self.training_worker.output_signal.connect(self.console.append)
self.training_worker.progress_signal.connect(self.progress_bar.setValue)
self.training_worker.finished_signal.connect(self.training_finished)
self.training_worker.start()
def stop_training(self):
if self.training_worker:
self.training_worker.stop()
self.console.append("⏹ 用户停止训练")
def training_finished(self, return_code, message):
self.train_btn.setEnabled(True)
self.stop_btn.setEnabled(False)
self.progress_bar.setVisible(False)
if return_code == 0:
self.status_bar.showMessage("✅ 训练完成!")
self.console.append_success(message)
self.load_training_result("results.png")
self.tabs.setCurrentIndex(3)
else:
self.status_bar.showMessage("❌ 训练失败")
self.console.append_error(message)
self.console.append_error("请检查控制台输出获取详细错误信息")
# 提供具体建议
if "closed file" in message or "I/O" in message:
self.console.append_info("💡 文件I/O错误建议:")
self.console.append_info(" 1. 检查磁盘空间是否充足")
self.console.append_info(" 2. 以管理员身份运行程序")
self.console.append_info(" 3. 关闭杀毒软件或添加排除项")
self.console.append_info(" 4. 使用 '工具' -> '修复数据集路径'")
self.console.append_info(" 5. 减少 batch size 再试")
self.training_worker = None
def run_inference(self):
if not self.model_weights.text():
QMessageBox.warning(self, "警告", "请选择模型权重")
return
if not os.path.exists(self.model_weights.text()):
QMessageBox.warning(self, "警告", "模型权重文件不存在")
return
if not self.source_path.text():
QMessageBox.warning(self, "警告", "请选择源")
return
source = self.source_path.text()
conf = self.confidence.value()
source_type = self.source_type.currentText()
if source_type != "webcam" and not os.path.exists(source):
QMessageBox.warning(self, "警告", "源文件/目录不存在")
return
cmd = f"yolo task=detect mode=predict model={self.model_weights.text()} conf={conf} source={source} save={str(self.save_results_check.isChecked()).lower()}"
if source_type == "webcam":
cmd = f"yolo task=detect mode=predict model={self.model_weights.text()} conf={conf} source=0 save={str(self.save_results_check.isChecked()).lower()}"
self.console.append(f"🔍 运行推理:\\n{cmd}\\n{'-'*50}")
self.status_bar.showMessage("推理中…")
infer_worker = TrainingWorker(cmd, os.getcwd())
infer_worker.output_signal.connect(self.console.append)
infer_worker.finished_signal.connect(self.inference_finished)
infer_worker.start()
self.infer_btn.setEnabled(False)
def inference_finished(self, return_code, message):
self.infer_btn.setEnabled(True)
if return_code == 0:
self.status_bar.showMessage("✅ 推理完成!")
self.console.append_success(message)
predict_path = os.path.expanduser("~/runs/detect/predict")
if os.path.exists(predict_path):
images = glob.glob(os.path.join(predict_path, "*.jpg"))
if images:
latest = max(images, key=os.path.getctime)
self.inference_viewer.set_image(latest)
self.inference_info.setText(f"📸 {os.path.basename(latest)}")
else:
self.status_bar.showMessage("❌ 推理失败")
self.console.append_error(message)
def closeEvent(self, event):
if self.training_worker and self.training_worker.isRunning():
reply = QMessageBox.question(
self, '确认退出',
'训练正在进行中,确定要退出吗?',
QMessageBox.Yes | QMessageBox.No,
QMessageBox.No
)
if reply == QMessageBox.Yes:
self.training_worker.stop()
self.training_worker.wait()
event.accept()
else:
event.ignore()
else:
event.accept()
def main():
app = QApplication(sys.argv)
app.setStyle('Fusion')
app.setWindowIcon(QIcon())
window = YOLOFireDetectionGUI()
window.show()
sys.exit(app.exec_())
if __name__ == '__main__':
main()





