news 2026/9/11 1:22:31

PyTorch火焰识别CNN实战:灰边填充与轻量模型设计

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch火焰识别CNN实战:灰边填充与轻量模型设计

简介:本资源是一套基于PyTorch实现的火焰检测深度学习项目,面向计算机视觉初学者与AI实践者,解决工业监控、森林防火等场景下的火焰图像二分类识别问题。项目采用CNN架构,完整覆盖数据预处理、模型训练与可视化交互全流程,含数据增强(灰边填充正方形化、随机旋转)、标签文本生成、模型训练及PyQt5图形界面部署。压缩包共302个文件,主体为288张JPG与8张PNG火焰/非火焰实拍图,辅以3个核心Python脚本(数据集构建、模型训练、UI启动)及3个配置/标签文本,整体11.72MB,结构清晰、开箱即用。已有75人学习下载,提供可直接运行的代码框架、规范的数据组织方式、详细的环境配置指引(含免安装包链接),并内置实际图片样本(如1_flip.jpg等增强样本),便于理解数据扩增逻辑与模型输入适配机制。

1. 火焰识别不是“拍张照就出结果”,而是要让CNN在灰边、旋转、标签映射的细节里学会区分火与非火

很多刚接触工业视觉检测的人,看到“火焰识别”四个字,第一反应是找现成模型跑一张图——结果要么把打火机当火灾报警,要么漏掉暗光下的明火。这个基于 PyTorch 的 CNN 实现,恰恰反其道而行:它不依赖 ImageNet 预训练权重微调,而是从零构建一个轻量但鲁棒的二分类网络,专为火焰这类高动态、低信噪比、易受光照干扰的目标设计。项目核心不在模型结构多炫酷,而在数据预处理层就埋下判别力——比如强制将所有输入缩放到正方形时,不是简单裁剪或拉伸,而是沿短边补灰边(gray padding),保留原始宽高比和火焰形态完整性;再叠加随机旋转(±15°)、亮度扰动(0.8–1.2 倍)和水平翻转,让模型真正学到“火焰的热辐射纹理特征”,而非记住某几张图的背景色块。适合需要快速验证火焰检测 pipeline 的安防集成商、消防设备厂商嵌入式工程师,以及正在做毕业设计的计算机视觉方向本科生——你不需要 GPU 服务器,一块 RTX 3060 就能完成全流程训练与推理。


2. 数据集构建:灰边填充 + 标签映射 + TXT 划分,三步锁定可复现的数据流

2.1 为什么必须用灰边填充而非裁剪?——火焰区域常位于图像边缘

火焰在监控画面中往往出现在画面顶部(如仓库顶棚起火)或角落(如配电箱冒烟),简单中心裁剪会直接切掉关键区域。本项目采用cv2.copyMakeBorder实现灰边填充,逻辑清晰且可控:

import cv2 import numpy as np def pad_to_square(img_path, target_size=224): img = cv2.imread(img_path) h, w = img.shape[:2] top, bottom = 0, 0 left, right = 0, 0 if h > w: pad = (h - w) // 2 left, right = pad, h - w - pad else: pad = (w - h) // 2 top, bottom = pad, w - h - pad # 使用灰度值 128 填充(非黑非白,减少模型对极端亮度的过拟合) padded = cv2.copyMakeBorder(img, top, bottom, left, right, cv2.BORDER_CONSTANT, value=(128, 128, 128)) return cv2.resize(padded, (target_size, target_size)) # 示例:处理 26.jpg padded_img = pad_to_square("dataset/fire/26.jpg") print(f"Original shape: {cv2.imread('dataset/fire/26.jpg').shape}, Padded shape: {padded_img.shape}")

提示:灰边值设为(128, 128, 128)而非(0, 0, 0)(255, 255, 255),是因为纯黑/白边会成为模型的强线索——它可能学会“只要看到黑边就判非火”,而非真正理解火焰纹理。128 是 RGB 中性灰,在 HSV 空间中饱和度为 0,避免引入虚假颜色特征。

2.201数据集文本生成制作.py的真实作用:生成带绝对路径的 train/val 划分文件

该脚本并非仅生成train.txtval.txt,而是构建了完整的数据加载契约。它按如下逻辑执行:

  • 扫描dataset/fire/dataset/no_fire/两个文件夹;
  • 对每个图片路径,拼接绝对路径(避免相对路径在不同工作目录下失效);
  • 按 8:2 比例随机划分(种子固定为 42,保证可复现);
  • 每行格式为:/absolute/path/to/26.jpg 1(1 表示 fire 类,0 表示 no_fire);
  • 同时生成class_names.txt,内容为:
    no_fire fire

关键代码段(已去除非核心日志):

import os import random from pathlib import Path def generate_dataset_txt(data_root, train_ratio=0.8, seed=42): random.seed(seed) fire_dir = Path(data_root) / "fire" no_fire_dir = Path(data_root) / "no_fire" fire_list = [str(f) for f in fire_dir.glob("*.jpg")] no_fire_list = [str(f) for f in no_fire_dir.glob("*.jpg")] # 打乱并划分 random.shuffle(fire_list) random.shuffle(no_fire_list) fire_train = fire_list[:int(len(fire_list)*train_ratio)] fire_val = fire_list[int(len(fire_list)*train_ratio):] no_fire_train = no_fire_list[:int(len(no_fire_list)*train_ratio)] no_fire_val = no_fire_list[int(len(no_fire_list)*train_ratio):] # 写入 train.txt with open("train.txt", "w") as f: for p in fire_train: f.write(f"{p} 1\n") for p in no_fire_train: f.write(f"{p} 0\n") # 写入 val.txt with open("val.txt", "w") as f: for p in fire_val: f.write(f"{p} 1\n") for p in no_fire_val: f.write(f"{p} 0\n") # 写入 class_names.txt with open("class_names.txt", "w") as f: f.write("no_fire\nfire\n") generate_dataset_txt("dataset/")

注意train.txtval.txt必须与02深度学习模型训练.py在同一目录,否则torch.utils.data.Dataset初始化时会因路径错误抛出FileNotFoundError。若你调整了数据集位置,请同步修改此脚本中的data_root参数。

2.3 图像增强策略表:旋转、亮度、翻转参数的实际影响

增强操作PyTorch 实现方式参数范围对火焰识别的关键作用典型失败场景
随机旋转transforms.RandomRotation(degrees=(-15, 15))±15°模拟监控摄像头轻微抖动,防止模型对火焰朝向过拟合旋转后火焰被裁出边界 → 本项目先灰边再旋转,规避此问题
随机亮度transforms.ColorJitter(brightness=(0.8, 1.2))0.8–1.2 倍应对黄昏/夜间/强光反射等复杂光照,提升泛化性亮度=0.8 时暗火易被误判为阴影 → 需配合灰边保留结构
水平翻转transforms.RandomHorizontalFlip(p=0.5)概率 0.5增加样本多样性,尤其对左右不对称火焰有效翻转后火焰与背景融合度变化 → 训练时需足够 epoch 收敛

这些增强全部封装在torchvision.transforms.Compose中,直接作用于Dataset.__getitem__()返回的 PIL Image,无需手动保存增强后图片,节省磁盘空间。


3. CNN 模型构建与训练:从 ResNet18 改造到火焰专用轻量结构

3.1 为什么不用完整 ResNet18?——去掉最后两层全连接,替换为双层小网络

原始 ResNet18 输出 1000 维(ImageNet 类别数),而本任务只需 2 维(fire/no_fire)。若直接微调,末端大参数量会拖慢收敛且易过拟合小数据集(当前仅约 200 张图)。项目采用“冻结主干 + 替换头部”策略:

  • 冻结layer1layer4的所有 BatchNorm 层(model.layer1[0].bn1.training = False);
  • fc层替换为:
    self.classifier = nn.Sequential( nn.Linear(512, 128), # ResNet18 最后一层输出通道为 512 nn.ReLU(inplace=True), nn.Dropout(0.3), nn.Linear(128, 2) )

完整模型定义(models/cnn_flame.py):

import torch import torch.nn as nn from torchvision import models class FlameCNN(nn.Module): def __init__(self, num_classes=2): super(FlameCNN, self).__init__() # 加载预训练 ResNet18,但不加载 fc 层权重 backbone = models.resnet18(pretrained=True) self.features = nn.Sequential(*list(backbone.children())[:-1]) # 去掉 avgpool 和 fc # 冻结前 3 个 layer 的参数(只训练 layer4 和 classifier) for param in self.features[0].parameters(): # conv1 param.requires_grad = False for param in self.features[1].parameters(): # bn1 param.requires_grad = False for param in self.features[4].parameters(): # layer1 param.requires_grad = False for param in self.features[5].parameters(): # layer2 param.requires_grad = False self.classifier = nn.Sequential( nn.Linear(512, 128), nn.ReLU(inplace=True), nn.Dropout(0.3), nn.Linear(128, num_classes) ) def forward(self, x): x = self.features(x) # 输出 shape: [B, 512, 1, 1] x = torch.flatten(x, 1) # 展平为 [B, 512] x = self.classifier(x) return x model = FlameCNN() print(model)

逻辑说明list(backbone.children())[:-1]移除了原始 ResNet18 的全局平均池化层(AdaptiveAvgPool2d)和全连接层(Linear),保留了从conv1layer4的全部卷积特征提取能力。torch.flatten(x, 1)[B, 512, 1, 1]变为[B, 512],适配后续全连接层输入维度。Dropout 0.3 是针对小数据集防止过拟合的关键正则项。

3.202深度学习模型训练.py的核心训练循环:带早停与最佳模型保存

训练脚本不使用torch.optim.lr_scheduler.ReduceLROnPlateau,而是采用更稳定的 StepLR,并内置早停(patience=10):

from torch.optim import Adam from torch.optim.lr_scheduler import StepLR import torch.nn.functional as F criterion = nn.CrossEntropyLoss() optimizer = Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=0.001) scheduler = StepLR(optimizer, step_size=7, gamma=0.1) # 每 7 个 epoch 学习率 ×0.1 best_val_acc = 0.0 patience_counter = 0 for epoch in range(50): model.train() train_loss = 0.0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() train_loss += loss.item() # 验证 model.eval() val_correct = 0 val_total = 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = torch.max(outputs.data, 1) val_total += labels.size(0) val_correct += (predicted == labels).sum().item() val_acc = 100 * val_correct / val_total print(f'Epoch {epoch+1}, Train Loss: {train_loss/len(train_loader):.4f}, Val Acc: {val_acc:.2f}%') # 早停与保存 if val_acc > best_val_acc: best_val_acc = val_acc torch.save(model.state_dict(), "best_flame_model.pth") patience_counter = 0 else: patience_counter += 1 if patience_counter >= 10: print("Early stopping triggered.") break scheduler.step()

参数说明filter(lambda p: p.requires_grad, model.parameters())确保只优化layer4classifier的参数,冻结部分大幅降低显存占用(RTX 3060 上 batch_size 可设为 32)。StepLRstep_size=7gamma=0.1意味着每 7 个 epoch 学习率衰减为原来的 1/10,避免后期震荡。


4. PyQt5 UI 集成:从单图推理到实时视频流,封装成可交付的检测界面

4.103pyqt_ui界面.py的三层架构:信号槽驱动 + OpenCV 解码 + 模型推理流水线

UI 不是简单弹窗,而是实现了“选择图片→点击检测→显示结果+置信度”的闭环。核心在于QThread封装推理过程,避免 GUI 卡死:

from PyQt5.QtCore import QThread, pyqtSignal import cv2 import torch from torchvision import transforms class InferenceThread(QThread): result_signal = pyqtSignal(str, float) # label, confidence def __init__(self, model_path, image_path): super().__init__() self.model_path = model_path self.image_path = image_path def run(self): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = FlameCNN() model.load_state_dict(torch.load(self.model_path, map_location=device)) model.to(device).eval() # 图像预处理(与训练一致) transform = transforms.Compose([ transforms.ToPILImage(), transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) img = cv2.imread(self.image_path) img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) tensor_img = transform(img_rgb).unsqueeze(0).to(device) # [1, 3, 224, 224] with torch.no_grad(): outputs = model(tensor_img) probs = F.softmax(outputs, dim=1) confidence, pred_class = torch.max(probs, 1) label_map = {0: "NO FIRE", 1: "FIRE"} self.result_signal.emit(label_map[pred_class.item()], confidence.item()) # 在主窗口中调用 def on_detect_clicked(self): if not self.current_image_path: return self.thread = InferenceThread("best_flame_model.pth", self.current_image_path) self.thread.result_signal.connect(self.show_result) self.thread.start()

逻辑说明transforms.Normalize使用 ImageNet 的均值标准差,是因为 ResNet18 主干是在该分布上预训练的,保持输入分布一致才能激活迁移学习效果。unsqueeze(0)添加 batch 维度,使输入 shape 符合模型要求[1, 3, 224, 224]F.softmax将 logits 转为概率,confidence.item()提取标量置信度供 UI 显示。

4.2 视频流检测的隐藏技巧:跳帧 + ROI 裁剪 + 置信度阈值过滤

若需接入 IPCAM 视频流,直接逐帧推理会导致延迟过高。项目预留了扩展接口,实际部署时建议:

  • 跳帧:每 3 帧处理 1 帧(cap.set(cv2.CAP_PROP_POS_FRAMES, frame_id));
  • ROI 裁剪:只检测画面中央 50% 区域(frame[h//4:3*h//4, w//4:3*w//4]),避开无关背景;
  • 置信度阈值if confidence.item() > 0.85:才触发报警,避免低置信度抖动。
# 视频流推理伪代码(需在 QTimer 中调用) def process_video_frame(self): ret, frame = self.cap.read() if not ret: return # ROI 裁剪 h, w = frame.shape[:2] roi = frame[h//4:3*h//4, w//4:3*w//4] # 推理(同单图流程,略) # ... if pred_label == "FIRE" and confidence > 0.85: cv2.putText(frame, "ALERT: FIRE!", (50, 50), cv2.FONT_HERSHEY_SIMPLEX, 1.2, (0, 0, 255), 2)

注意:PyQt5 的QLabel.setPixmap()无法直接显示 OpenCV 的 BGR 图像,必须先转换:qimg = QImage(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB), w, h, w*3, QImage.Format_RGB888)


5. 模型验证与边界 case 处理:用混淆矩阵定位漏检/误报根源

5.1 生成验证报告:不只是准确率,更要看出哪类样本在拖后腿

运行val_report.py(需自行编写,但逻辑极简)可输出详细分类报告:

from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns import matplotlib.pyplot as plt # 假设已有 all_labels 和 all_preds 列表 print(classification_report(all_labels, all_preds, target_names=["no_fire", "fire"])) cm = confusion_matrix(all_labels, all_preds) sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", xticklabels=["no_fire", "fire"], yticklabels=["no_fire", "fire"]) plt.ylabel("True Label") plt.xlabel("Predicted Label") plt.title("Confusion Matrix") plt.savefig("confusion_matrix.png", dpi=300, bbox_inches='tight')

典型输出解读:

precision recall f1-score support no_fire 0.92 0.96 0.94 50 fire 0.96 0.92 0.94 50 accuracy 0.94 100
  • fire类 recall=0.92 表示 100 张火图中有 8 张漏检;
  • no_fire类 precision 仅 0.75,则说明模型把 25 张非火图错判为火——此时应检查no_fire类中是否混入了暖色灯光、夕阳、红色广告牌等干扰样本,并针对性扩充 negative 数据。

5.2 三个必测边界 case 及应对方案

边界 case测试方法预期表现修复动作
暗光火焰(如夜间仓库)用手机拍摄暗处打火机火焰,亮度调至最低置信度 < 0.6 → 模型未学好低对比度特征01数据集文本生成制作.py中增加transforms.ColorJitter(brightness=(0.4, 1.0))并重训
小火焰(如电线短路火花)将 17.jpg 缩放至 32×32 后测试分类错误 → 模型感受野不足修改pad_to_squaretarget_size=256,增大输入分辨率
火焰+烟雾遮挡(如初期火灾)合成烟雾图层叠加到火图上误判为 no_fire → 特征被烟雾抑制在训练时加入transforms.GaussianBlur(kernel_size=(3,3), sigma=(0.1, 2.0))模拟烟雾模糊

提示:所有修复动作都应在01数据集文本生成制作.pytransforms.Compose中统一调整,确保训练/验证/推理预处理完全一致。切勿在推理时单独改 transform,否则导致线上效果劣于离线测试。

验证完成后,将best_flame_model.pthclass_names.txt一并打包,即可交付给嵌入式团队部署到 Jetson Nano 或 RK3399 等边缘设备——此时模型体积约 42MB,FP16 量化后可压至 22MB,满足工业级实时性要求。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 1:21:25

AI时代产品经理价值重构:从需求搬运工到AI Agent架构师

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 1:18:13

New Concept English: Practice Progress

New Concept English: Practice & Progress1. 新概念精讲 - DianaN. Practice & ProgressReferences1. 新概念精讲 - Diana 01-20 https://www.youtube.com/playlist?listPLEPELt2_eXcqMicCQVKO92pl86-z4UFMV 21-40 https://www.youtube.com/playlist?listPLEPELt2…

作者头像 李华
网站建设 2026/9/11 1:07:50

KMP算法详解:从暴力匹配到next数组的完整推导与实现

写 KMP 算法这篇文章&#xff0c;其实是我早就想做的事。字符串匹配是写代码几乎绕不开的一件事&#xff0c;不管你是刷 LeetCode、打信奥、做文本处理&#xff0c;还是写搜索引擎、做日志分析&#xff0c;KMP&#xff08;Knuth-Morris-Pratt&#xff09;算法都是绕不过去的一道…

作者头像 李华