PyTorch ResNet真假图片识别实战:从训练到PyQt界面部署
发布时间:2026/10/1 18:00:32来源:尧图网络
简介这份资源面向具备Python与PyTorch基础、希望入门图像真伪识别实战的开发者与学习者提供一套基于ResNet与CNN训练识别真假图片的完整代码方案。压缩包共7个文件约190KB包含3个py脚本、2张示例提示图、1个txt依赖清单和1份docx说明文档分别承担数据路径生成、模型训练与界面演示等用途。代码不含数据集图片需自行搜集图片放入对应分类文件夹每个文件夹内附有提示图指引存放位置便于快速组织数据。运行数据路径生成脚本可自动划分训练集与验证集并输出txt标签文件训练脚本会读取该文件进行训练并适配分类文件夹数量变化新增类别无需改动代码训练过程带进度条每个epoch显示准确率与损失值结束后保存日志与模型权重。目前已有46人学习适合想掌握CNN图像分类流程、理解数据组织与训练监控细节的读者参考。1. 拆开这个 ResNet 真假图片识别包三个 py 文件能跑通什么上周有个做电商的朋友找我说他们审核团队每天要人工过几千张商品图想用模型先把明显是 AI 生成或者盗图的筛出来。我第一反应是找现成的开源方案翻了一圈发现大部分要么依赖特定数据集要么代码封装得太深改不动。后来拆了这个 ResNet 真假图片识别包三个 py 文件加一份说明文档结构简单到有点意外——但恰恰是这种简单让它适合拿来当二次开发的底子。这个资源的核心是用 PyTorch 搭一个 ResNet 分类器通过 CNN 训练来区分真图和假图。它不含数据集图片需要你自己往对应文件夹里放图每个文件夹里有一张提示图告诉你放哪。代码适配了分类文件夹个数你加新类别不用改训练脚本。适合谁有 Python 和 PyTorch 基础、想快速搭一个图像二分类或小多分类流程的从业者。如果你连环境都没装过建议先看说明文档里的环境配置部分或者参考博文把 PyTorch 装好再动手。2. 环境搭建与数据准备从 requirement.txt 到文件夹结构2.1 依赖安装与版本选择拿到包先看 requirement.txt这是最省事的入口。我一般不会直接pip install -r而是先看一眼里面有没有版本锁死导致冲突的包。这个项目的依赖比较常规核心就是 torch、torchvision、numpy、Pillow可能还有 tqdm 用来显示进度条。# 建议先建虚拟环境避免和系统里的包打架 python -m venv venv_resnet # Windows 下激活 venv_resnet\Scripts\activate # Linux/Mac 下激活 source venv_resnet/bin/activate # 安装依赖如果 requirement.txt 里没锁版本可以手动指定稳定版本 pip install torch torchvision numpy Pillow tqdm逻辑说明虚拟环境是底线操作我见过太多人因为全局环境里 torch 版本和 CUDA 不匹配训练时直接报RuntimeError: CUDA error。参数方面torch 版本建议选 1.10 以上torchvision 对应版本即可。如果你没有 NVIDIA 显卡装 CPU 版也能跑只是训练慢小数据集二分类大概每 epoch 几分钟能接受。提示安装完成后用python -c import torch; print(torch.__version__, torch.cuda.is_available())验证一下输出 True 说明 GPU 可用。2.2 数据集文件夹的摆放规则这个包不含图片所以第一步是收集数据。它的设计逻辑是按文件夹名分类比如你建两个文件夹真图片和假图片每个文件夹里放对应类别的 jpg 或 png。每个文件夹里有一张提示图告诉你图片应该放在这一层不要嵌套子文件夹。我一般会按 8:2 的比例手动分一下训练和验证但这个脚本是在01生成txt.py里自动划分的所以你只需要把所有图片混在一起放脚本会按比例切分。常见做法是每个类别至少准备 200 张以上太少的话 ResNet 容易过拟合验证准确率会虚高。# 目录结构示例 dataset/ ├── 真图片/ │ ├── 1.jpg │ ├── 2.jpg │ └── ...提示图也在里面脚本会跳过或你手动删掉 └── 假图片/ ├── 1.jpg ├── 2.jpg └── ...参数说明图片格式建议统一成 jpg尺寸不用提前裁剪脚本里一般会做 resize。如果你收集的图片长宽比差异很大可以在训练脚本里改一下Resize和CenterCrop的参数这个后面讲。2.3 生成训练索引01生成txt.py 逐行拆解这个脚本的作用是遍历数据集文件夹把图片路径和对应标签写进 txt同时划分训练集和验证集。我拆过不少类似脚本这个写得比较直白适合改。import os import random # 数据集根目录改成你自己的路径 data_root ./dataset # 输出 txt 的路径 train_txt ./train.txt val_txt ./val.txt # 验证集比例 val_ratio 0.2 # 获取所有类别文件夹名按名称排序保证标签一致 classes sorted(os.listdir(data_root)) class_to_idx {cls: idx for idx, cls in enumerate(classes)} all_samples [] for cls in classes: cls_dir os.path.join(data_root, cls) if not os.path.isdir(cls_dir): continue for img_name in os.listdir(cls_dir): # 跳过提示图假设提示图命名包含 提示 或 example if 提示 in img_name or example in img_name.lower(): continue img_path os.path.join(cls_dir, img_name) all_samples.append((img_path, class_to_idx[cls])) # 随机打乱后按比例划分 random.shuffle(all_samples) split int(len(all_samples) * (1 - val_ratio)) train_samples all_samples[:split] val_samples all_samples[split:] # 写入 txt格式路径 标签 with open(train_txt, w, encodingutf-8) as f: for path, label in train_samples: f.write(f{path} {label}\n) with open(val_txt, w, encodingutf-8) as f: for path, label in val_samples: f.write(f{path} {label}\n) print(f训练集 {len(train_samples)} 张验证集 {len(val_samples)} 张)逻辑说明先按文件夹名排序生成类别索引这样即使你后面加了新文件夹只要排序位置不变标签就不会乱。跳过提示图的逻辑是硬编码的如果你的提示图命名不一样需要自己改一下判断条件。参数方面val_ratio控制验证集比例二分类任务 0.2 够用多分类可以调到 0.3。注意如果你后面增加了类别文件夹重新运行这个脚本会重新生成 txt标签索引也会更新所以训练脚本必须每次重新读取 txt不能缓存旧标签。3. CNN 训练脚本02CNN训练数据集.py 的参数与训练循环3.1 模型选型为什么用 ResNet 而不是自己搭 CNN这个包用的是 ResNet具体是 resnet18 还是 resnet50 要看代码里的torchvision.models调用。我一般会先看它有没有加载预训练权重。如果加载了pretrainedTrue那在小数据集上收敛会快很多这是 ResNet 相比自己搭三层 CNN 的最大优势——残差连接让梯度能传得更深预训练权重提供了通用的特征提取能力。import torch import torch.nn as nn from torchvision import models # 加载 ResNet18如果不需要预训练权重可以设 pretrainedFalse model models.resnet18(pretrainedTrue) # 替换最后的全连接层适配你的类别数 num_classes 2 # 真/假 二分类 model.fc nn.Linear(model.fc.in_features, num_classes) # 如果有 GPU 就放到 GPU 上 device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device)逻辑说明model.fc.in_features是 ResNet 最后一层全连接层的输入维度resnet18 是 512resnet50 是 2048。替换成你的类别数后前面的卷积层参数要么冻结要么微调。常见做法是先用预训练权重冻结卷积层只训练全连接层几个 epoch再解冻全部微调。这个包如果没做冻结直接全量微调也能跑只是小数据集上容易过拟合。参数说明pretrainedTrue会下载预训练权重第一次运行需要联网。如果你环境不能联网需要提前把权重文件放到~/.cache/torch/hub/checkpoints/下。3.2 数据加载与增强Dataset 和 DataLoader 的配置训练脚本里一般会自定义一个 Dataset 类读取 txt 里的路径和标签然后用 DataLoader 批量加载。数据增强这块训练集通常做 RandomResizedCrop、RandomHorizontalFlip验证集只做 Resize 和 CenterCrop。from torch.utils.data import Dataset, DataLoader from torchvision import transforms from PIL import Image class ImageDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] with open(txt_path, r, encodingutf-8) as f: for line in f: path, label line.strip().split() self.samples.append((path, int(label))) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label # 训练集增强 train_transform transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 验证集只做 resize 和归一化 val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_dataset ImageDataset(./train.txt, transformtrain_transform) val_dataset ImageDataset(./val.txt, transformval_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4)逻辑说明Normalize的均值和方差是 ImageNet 的统计值因为用了预训练权重必须保持一致。batch_size根据显存调8G 显存跑 resnet18 可以到 64跑 resnet50 建议 32 以下。num_workers在 Windows 下有时候会报错设成 0 可以规避。参数说明RandomResizedCrop(224)表示随机裁剪并缩放到 224x224这是 ResNet 的标准输入尺寸。如果你的图片普遍很小可以改成 128但需要同步改模型输入。3.3 训练循环与日志保存每个 epoch 看什么训练循环里一般会打印进度条、准确率和损失值。这个包用了 tqdm 的话你能看到大概还要多久。每个 epoch 结束后会保存 log 日志记录准确率和损失方便后面画曲线。import torch.optim as optim from tqdm import tqdm criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-4) num_epochs 20 best_acc 0.0 for epoch in range(num_epochs): model.train() running_loss 0.0 correct 0 total 0 # 训练阶段带进度条 train_bar tqdm(train_loader, descfEpoch {epoch1}/{num_epochs}) for imgs, labels in train_bar: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() train_bar.set_postfix(lossloss.item(), acccorrect/total) train_acc correct / total train_loss running_loss / len(train_loader) # 验证阶段 model.eval() val_correct 0 val_total 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) _, predicted torch.max(outputs, 1) val_total labels.size(0) val_correct (predicted labels).sum().item() val_acc val_correct / val_total print(fEpoch {epoch1}: train_loss{train_loss:.4f}, train_acc{train_acc:.4f}, val_acc{val_acc:.4f}) # 保存 log with open(train_log.txt, a, encodingutf-8) as f: f.write(f{epoch1},{train_loss:.4f},{train_acc:.4f},{val_acc:.4f}\n) # 保存最佳模型 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) print(f模型已保存当前最佳验证准确率 {best_acc:.4f})逻辑说明optimizer.zero_grad()必须在反向传播前调用否则梯度会累加。model.eval()和torch.no_grad()在验证阶段必须加否则 BN 层和 Dropout 会继续更新验证结果不准。参数方面学习率lr1e-4是微调预训练模型的常见起点如果 loss 下降很慢可以调到 1e-3如果震荡就降到 1e-5。注意log 文件是追加写入的如果你重新训练记得先删掉旧的 log否则曲线会混在一起。4. PyQt 界面与推理03pyqt界面.py 怎么把模型用起来4.1 界面逻辑与模型加载这个包带了 PyQt 界面说明作者考虑到了非技术用户的使用场景。界面一般是一个窗口一个按钮选择图片一个标签显示预测结果。模型加载部分要注意必须和训练时的模型结构完全一致否则load_state_dict会报 key 不匹配。import sys from PyQt5.QtWidgets import QApplication, QMainWindow, QLabel, QPushButton, QVBoxLayout, QWidget, QFileDialog from PyQt5.QtGui import QPixmap import torch from torchvision import models, transforms from PIL import Image class MainWindow(QMainWindow): def __init__(self): super().__init__() self.setWindowTitle(真假图片识别) self.resize(400, 500) # 加载模型 self.model models.resnet18(pretrainedFalse) self.model.fc torch.nn.Linear(self.model.fc.in_features, 2) self.model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) self.model.eval() self.transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 界面组件 self.label QLabel(请选择图片) self.btn QPushButton(选择图片) self.btn.clicked.connect(self.open_image) layout QVBoxLayout() layout.addWidget(self.label) layout.addWidget(self.btn) container QWidget() container.setLayout(layout) self.setCentralWidget(container) def open_image(self): path, _ QFileDialog.getOpenFileName(self, 选择图片, , Images (*.jpg *.png)) if not path: return pixmap QPixmap(path) self.label.setPixmap(pixmap.scaled(300, 300)) img Image.open(path).convert(RGB) img_tensor self.transform(img).unsqueeze(0) with torch.no_grad(): output self.model(img_tensor) _, pred torch.max(output, 1) self.label.setText(f预测结果{真图片 if pred.item() 0 else 假图片}) app QApplication(sys.argv) window MainWindow() window.show() sys.exit(app.exec_())逻辑说明map_locationcpu是为了在没有 GPU 的机器上也能加载。unsqueeze(0)是给图片加一个 batch 维度因为模型输入要求是 4 维。参数方面类别索引要和训练时一致训练时真图片排在前面对应 0假图片对应 1界面里的判断要对应上。4.2 推理速度与批量预测单张推理在 CPU 上大概几十毫秒GPU 上几毫秒。如果你要批量筛图不建议走界面直接写个脚本遍历文件夹更高效。import os from PIL import Image def batch_predict(folder, model, transform, device): results [] for img_name in os.listdir(folder): img_path os.path.join(folder, img_name) img Image.open(img_path).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): output model(img_tensor) _, pred torch.max(output, 1) results.append((img_name, pred.item())) return results逻辑说明批量预测时model.eval()和torch.no_grad()必须加否则显存会爆。参数方面如果图片特别多可以每 100 张清一次缓存torch.cuda.empty_cache()。5. 避坑与排查训练不收敛、显存溢出、标签错位5.1 损失值不下降准确率卡在 50%现象二分类任务训练几个 epoch 后loss 一直在 0.69 附近准确率 50% 左右和随机猜一样。原因最常见的是标签和图片没对应上比如 txt 里路径写错了或者类别索引乱了。另一个可能是学习率太大模型直接震荡。解决先检查 train.txt 里的路径能不能打开用Image.open试几张。然后打印一下每个 batch 的标签分布确认不是全 0 或全 1。学习率降到 1e-5 再试。5.2 CUDA out of memory现象训练开始几秒后报RuntimeError: CUDA out of memory。原因batch_size 太大或者图片分辨率太高或者没有用torch.no_grad()导致验证阶段显存累积。解决把 batch_size 减半图片 resize 到 128验证阶段加上with torch.no_grad():。如果还不行换 resnet18 或者用 CPU 跑。5.3 验证准确率远高于训练准确率现象训练准确率 70%验证准确率 95%。原因验证集太小或者验证集和训练集有重叠图片。这个脚本是随机划分的如果你图片本身有重复就会泄漏。解决检查验证集图片是否在训练集里出现过去重后再跑。另外验证集至少要有 50 张以上否则指标波动很大。5.4 PyQt 界面报 no Qt platform plugin现象运行03pyqt界面.py时报This application failed to start because no Qt platform plugin could be initialized。原因PyQt5 的环境变量没配好或者 conda 环境里混了多个 Qt 版本。解决设置环境变量export QT_QPA_PLATFORM_PLUGIN_PATH$VIRTUAL_ENV/lib/python3.x/site-packages/PyQt5/Qt5/plugins/platformsWindows 下在系统环境变量里加。或者直接pip install pyqt5 --force-reinstall。5.5 增加类别后训练报错现象加了第三个文件夹后训练时CrossEntropyLoss报维度不匹配。原因模型最后的全连接层还是 2 输出但标签有 0、1、2 三个值。解决在训练脚本里把num_classes改成len(classes)重新运行01生成txt.py和训练脚本。这个包号称适配分类文件夹个数但前提是你改了num_classes这个变量或者代码里自动读取了类别数。6. 进阶技巧用混淆矩阵和阈值调优把误判压下去训练完模型准确率只是一个粗指标。真假图片识别这种场景误判的代价不对称——把真图判成假图可能只是多一次人工复核但把假图判成真图可能直接放行风险内容。所以我一般会看混淆矩阵然后调分类阈值。from sklearn.metrics import confusion_matrix, classification_report import numpy as np model.eval() all_preds [] all_labels [] with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) outputs model(imgs) # 取 softmax 后的概率 probs torch.softmax(outputs, dim1) _, preds torch.max(probs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) # 混淆矩阵 cm confusion_matrix(all_labels, all_preds) print(混淆矩阵) print(cm) print(classification_report(all_labels, all_preds, target_names[真图片, 假图片]))逻辑说明classification_report会输出每个类别的 precision、recall、f1-score。如果假图片的 recall 偏低说明漏检多需要把判定为假图片的阈值降低。默认是取最大概率的类别你可以改成概率大于 0.4 就判为假图片。# 自定义阈值假图片概率大于 0.4 就判为假 threshold 0.4 probs torch.softmax(outputs, dim1) preds (probs[:, 1] threshold).long()参数说明阈值调低会提高假图片的召回率但会牺牲精确率也就是更多真图被误判。具体调到多少看你的业务能接受多少误杀。我一般会画一个 precision-recall 曲线找平衡点。还有一个技巧是测试时增强TTA对同一张图做多次裁剪和翻转取平均概率。这个在验证集上能涨 1-2 个点但推理时间翻倍。def tta_predict(model, img, transform_list): probs [] for t in transform_list: img_tensor t(img).unsqueeze(0).to(device) with torch.no_grad(): output model(img_tensor) prob torch.softmax(output, dim1) probs.append(prob.cpu().numpy()) return np.mean(probs, axis0)从那以后我每次训练完分类模型都强制走一遍混淆矩阵和阈值扫描不再只看一个准确率数字。这个包的结构简单正好适合拿来练手这套流程。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网