拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch+Flask+PyQt实现ResNet动物图像分类系统全链路

简介基于ResNet的动物图像分类系统完整源码面向Python初学者及期末大作业、课程设计场景融合PyQt、Flask、HTML5与PyTorch技术栈覆盖数据处理、模型训练、推理预测到多端部署全流程。资源共27个文件压缩包大小41.75MB包含8个Python源码如train.py、predict.py、myflask.py、window.py等对应训练、预测、Flask后端与PyQt界面、11张PNG图片界面及效果预览、1个PyTorch权重文件、HTML模板及README说明文档结构完整。已有333人学习下载。源码附有详细注释提供预训练模型权重与测试图片简单部署即可运行既可作为课程设计/期末大作业直接提交也能帮助新手轻松理解ResNet图像分类与前后端交互的实战流程。1. 从期末大作业到可演示系统ResNet 动物分类的完整链路一个用 PyTorch 训练出的 ResNet 模型如果只能打印训练日志它只是一段脚本如果还能同时被网页和桌面程序调用它就变成了一个完整的动物图像分类系统。这个项目的典型架构是PyTorch 负责模型的训练与推理Flask 把模型包装成 HTTP 接口HTML5 页面和 PyQt 窗口分别充当两种客户端三端共用同一份源代码和文档说明。我见过不少期末大作业只停在训练脚本模型调通了、准确率打出来了但老师运行时还得靠命令行传参演示效果大打折扣。把 Flask 和 PyQt 都接上等于把数据预处理、模型微调、服务端接口、前端交互四条链路一次性展示出来答辩时能讲的内容多一倍代码量却只多了两个文件。这篇文章按这条链路把每一段的代码和参数讲清楚照着抄就能跑通。适合照着做的人群有两类一类是正在为 Python 期末大作业找完整方案的学生另一类是想快速确认训练到部署最小链路怎么串的工程师。前者照搬整体结构后者可以直接只取 Flask 和 PyTorch 两段。2. 用 ResNet 预训练模型搭动物分类器数据集整理与 PyTorch 微调把 ResNet 用在动物分类上第一个要做的决定是从头训练还是微调。对期末作业的数据量通常几百到几千张图来说从头训练一个 ResNet 几乎不可能收敛到可用的精度因为 ResNet 的参数量远大于你的数据量。常见做法是加载在 ImageNet 上训练好的 resnet 预训练模型冻结大部分层只训练最后的全连接层这样几十个 epoch 就能拿到不错的效果。动手之前先确认环境。用 Anaconda 创建虚拟环境然后 pip 安装 PyTorchCPU 版本一条命令即可GPU 版本需要根据 CUDA 版本选择对应的安装组合。国内网络环境下建议加上清华镜像源pip install torch torchvision -i https://pypi.tuna.tsinghua.edu.cn/simple能避免下载超时的问题。PyTorch 基础框架里加载 ResNet 最省事的入口是torchvision.models下面按数据到模型的顺序拆开讲。2.1 数据集目录结构与 8:2 划分脚本torchvision.datasets.ImageFolder对目录结构有硬性要求每个类别一个文件夹文件夹名就是标签名图片放在类别文件夹里。这是整个项目里最容易省事也最容易出错的地方。组织好的目录长这样data/ train/ cat/001.jpg dog/001.jpg bird/001.jpg val/ cat/001.jpg dog/001.jpg bird/001.jpg实际的素材往往混在一个大目录里需要先做一次划分。下面的脚本按 8:2 比例把每类图片随机分到 train 和 valimport os import random import shutil src animals_raw # 原始素材根目录每类一个子文件夹 dst data random.seed(42) # 固定随机种子保证每次划分结果一致 for cls in os.listdir(src): cls_path os.path.join(src, cls) if not os.path.isdir(cls_path): continue imgs os.listdir(cls_path) random.shuffle(imgs) split int(len(imgs) * 0.8) for part, subset in [(train, imgs[:split]), (val, imgs[split:])]: out_dir os.path.join(dst, part, cls) os.makedirs(out_dir, exist_okTrue) for name in subset: shutil.copy(os.path.join(cls_path, name), os.path.join(out_dir, name)) print(f{cls}: train{split}, val{len(imgs) - split})划分的逻辑很简单先算出 80% 的索引作为训练集剩下的作为验证集。random.seed(42)是关键否则每次跑出来的划分都不一样排查问题时无法复现。注意shutil.copy会复制文件素材量大时可以换成os.rename直接移动省一半磁盘空间。类别文件夹的名字最终会变成分类结果里的标签建议直接用英文cat、dog、bird省去中文编码问题。2.2 transforms 与 DataLoader 参数数据加载决定了训练能不能收敛。ResNet 的标准输入是 224x224 的 RGB 图像并且要用 ImageNet 的均值方差做归一化。训练集和验证集的 transform 必须分开写训练集可以加随机增强验证集只做缩放和归一化from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224), # 随机裁剪并缩放到 224 transforms.RandomHorizontalFlip(), # 随机水平翻转相当于免费扩增数据 transforms.ToTensor(), # 转为张量像素归一化到 [0,1] transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), # 先缩放到 256 transforms.CenterCrop(224), # 再中心裁剪到 224 transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])Normalize里的均值和方差是 ImageNet 数据集的统计值对 resnet 预训练模型来说是固定搭配不要改成自己的数据集统计值否则预训练权重的作用会大打折扣。RandomResizedCrop的作用不只是缩放它通过随机裁剪让模型看到同一只动物的不同构图对动物这种主体位置不固定的图像特别有效。DataLoader 的参数直接影响训练速度和显存占用参数常用值说明batch_size16 / 32显存不够时优先降到 8而不是改模型结构shuffle只在训练集为 True验证集必须 False保证评估稳定num_workers2 / 4Windows 下建议设 0 或 2过高容易报错pin_memoryTrueGPU 训练时开启能少量提速from torchvision import datasets from torch.utils.data import DataLoader train_ds datasets.ImageFolder(data/train, transformtrain_transform) val_ds datasets.ImageFolder(data/val, transformval_transform) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers2, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers2, pin_memoryTrue) classes train_ds.classes # 按文件夹名字母序排列的标签列表 print(classes) # [bird, cat, dog]train_ds.classes是 ImageFolder 自动生成的标签列表按文件夹名的字母序排列后续推理时的索引到类别名的映射必须用它不能自己写死顺序。2.3 加载 resnet 预训练模型并替换全连接层模型加载是整套代码里最核心的 3 行。PyTorch 新版推荐用weights参数而不是pretrainedTrue后者已经标记为废弃import torch import torchvision.models as models model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 读取原模型的 fc 层输入维度替换成自己的类别数 num_ftrs model.fc.in_features model.fc torch.nn.Linear(num_ftrs, len(classes)) # 先冻结全部参数只放开最后一层 for param in model.parameters(): param.requires_grad False for param in model.fc.parameters(): param.requires_grad Trueresnet18 对期末作业的数据量是够用的单张 CPU 推理不到 100ms训练也快。如果追求更高的准确率可以换成 resnet50代价是训练时间翻倍、显存占用更高。model.fc.in_features是读取 ResNet 最后一个全连接层的输入维度ResNet18 是 512ResNet50 是 2048这个写法保证换模型时不用改代码。冻结策略要看数据量灵活调整。数据量只有几百张时只训练 fc 层最稳妥数据量超过两千张可以把最后两个残差块也解冻学习率设小一个量级效果会更好。解冻的写法是把对应模块的requires_grad重新设为 True其余不动。提示第一次跑通时先别追求准确率用 10 个类别、每个类别 50 张图把流程走通再回到这一步调数据量和模型规模。3. 训练脚本中的关键参数优化器、损失函数与模型保存时机训练阶段是跑得通和跑得好的分水岭。很多作业代码能跑但准确率一直上不去问题几乎都出在三个地方优化器选错、损失函数没有匹配任务、模型保存时机不对。这一章直接给出一份可以直接用的训练脚本然后逐个参数解释为什么这么设。3.1 训练循环与验证逻辑完整的训练脚本包含标准的三段式训练、验证、保存。下面这份代码可以直接放进train.py运行import torch import torch.nn as nn from torch.optim import AdamW device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() # 多分类任务的标准损失 optimizer AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr1e-4, weight_decay1e-2) def evaluate(model, loader): model.eval() correct total 0 with torch.no_grad(): # 验证阶段不计算梯度 for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) correct (outputs.argmax(1) labels).sum().item() total labels.size(0) return correct / total best_acc 0.0 for epoch in range(20): model.train() total_loss total_num 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() # 更新参数 total_loss loss.item() * images.size(0) total_num labels.size(0) val_acc evaluate(model, val_loader) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_animal_resnet.pth) print(fepoch{epoch1} loss{total_loss/total_num:.4f} fval_acc{val_acc:.4f} best{best_acc:.4f})这个循环里有四个点不能省optimizer.zero_grad()必须在loss.backward()之前否则梯度会在每个 batch 间累加model.train()和model.eval()不能混用因为 BatchNorm 和 Dropout 在两种模式下行为不同验证阶段包在torch.no_grad()里否则显存会被多余的计算图吃掉保存模型要用验证集准确率判断不能用训练集损失。损失函数在多分类任务里固定选CrossEntropyLoss它内部已经包含了 softmax所以模型最后的 fc 层输出不需要再手动过 softmax。优化器选 AdamW 而不是 Adamweight_decay1e-2的正则化对微调预训练模型有明显帮助能抑制过拟合。3.2 模型保存的两种方式与加载陷阱PyTorch 保存模型有两种常见姿势期末作业两种都能用但踩坑的方式不一样# 方式一只保存参数推荐文件小、跨环境兼容好 torch.save(model.state_dict(), best_animal_resnet.pth) # 方式二保存整个模型含结构但依赖原文件路径 torch.save(model, best_animal_resnet_full.pth)推荐用方式一。加载时先构建一遍模型结构再灌入参数def load_model(checkpoint_path, num_classes): model models.resnet18(weightsNone) # 只搭结构不加载预训练权重 model.fc torch.nn.Linear(model.fc.in_features, num_classes) state torch.load(checkpoint_path, map_locationcpu) model.load_state_dict(state) model.eval() return modelmap_locationcpu是关键。在 GPU 机器上训练的权重文件里记录的是 CUDA tensor换到没显卡的机器上直接torch.load会报错加上map_locationcpu就能在任何环境里加载。weightsNone只搭建 ResNet 结构不加载 ImageNet 预训练权重因为你要加载的是自己训练好的参数两者冲突。加载完一定要调用model.eval()否则推理结果可能是错的这个坑在模型保存后接入 Flask 时特别常见。3.3 学习率、batch_size、epochs 的搭配与过拟合信号微调场景下这几个参数有固定的搭配套路参数推荐区间场景说明lr1e-4 ~ 1e-5只训练 fc 层用 1e-4解冻更多层用 1e-5batch_size16 ~ 64越小梯度噪声越大越大越吃显存epochs15 ~ 30提前停止连续 5 个 epoch 验证集不涨就停类别数10 ~ 30期末作业的合理范围再多需要加大数据量训练时重点盯验证集准确率而不是训练集损失。如果训练集 loss 一直降、验证集准确率在某个值附近震荡就是过拟合信号优先做两件事加数据增强把RandomHorizontalFlip换成同时加RandomRotation(10)或者调大weight_decay到 3e-2。如果训练和验证的 loss 都不降说明学习率太大或数据没对齐先检查Normalize的均值方差和输入尺寸是不是 224。注意本机没有 GPU 就老老实实用 CPU 训练。ResNet18 32 张图一个 batchCPU 上一个 epoch 大概几分钟20 个 epoch 能接受。不要为了提速把num_workers调到 8Windows 上很容易触发多进程报错。4. 用 Flask 封装推理 API让 HTML5 页面直接识别动物模型训练好之后下一步是把best_animal_resnet.pth变成一个可以远程调用的服务。这里用 Flask 是因为它足够轻一个.py文件就能撑起完整的推理服务不需要像 Django 那样建项目。Flask 在 Python 生态里做模型部署接口是最常见的选择它的开发服务器对期末作业这种并发量极低的场景完全够用。4.1 Flask 应用结构模型预热与 /predict 接口推理服务有两个硬性要求模型只加载一次不能每个请求都 load图片预处理必须和训练时的验证集 transform 完全一致。下面是一个能直接运行的app.pyimport io from flask import Flask, request, jsonify from PIL import Image import torch import torchvision.transforms as transforms from model_zoo import load_model # 复用训练部分的加载函数 app Flask(__name__) app.config[MAX_CONTENT_LENGTH] 8 * 1024 * 1024 # 限制上传 8MB classes [bird, cat, dog] model load_model(best_animal_resnet.pth, len(classes)) # 模块级加载 infer_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) app.route(/predict, methods[POST]) def predict(): file request.files.get(image) if file is None: return jsonify({error: missing image}), 400 img Image.open(io.BytesIO(file.read())).convert(RGB) tensor infer_transform(img).unsqueeze(0) # 增加 batch 维 with torch.no_grad(): logits model(tensor) prob torch.softmax(logits, dim1)[0] top3_idx prob.topk(3).indices.tolist() result [{label: classes[i], prob: round(prob[i].item(), 4)} for i in top3_idx] return jsonify({top1: result[0], top3: result}) if __name__ __main__: app.run(host0.0.0.0, port5000, debugFalse)模型加载放在模块顶层进程启动时只执行一次后续请求直接复用。convert(RGB)是必要的兜底有些图片是 RGBA 或灰度模式不转成 RGB 会直接触发运行时错误。softmax把 fc 层输出转成概率topk(3)取概率最高的前三个类别前端就能展示最可能的三种动物而不是只给一个答案。4.2 HTML5 页面文件选择、fetch 上传与结果回显前端页面用一个static/index.html就能搞定。HTML5 的FormData可以直接把文件塞进 POST 请求配合fetch实现不刷新页面就能识别!DOCTYPE html html langzh-CN head meta charsetUTF-8 title动物图像分类系统/title /head body h2上传一张动物图片/h2 input typefile idfileInput acceptimage/* button onclickuploadImage()开始识别/button div idresult/div script async function uploadImage() { const fileInput document.getElementById(fileInput); if (!fileInput.files.length) { alert(请先选择图片); return; } const formData new FormData(); formData.append(image, fileInput.files[0]); const resp await fetch(/predict, { method: POST, body: formData }); const data await resp.json(); if (data.error) { document.getElementById(result).innerText 错误 data.error; return; } const r data.top1; document.getElementById(result).innerHTML 识别结果b${r.label}/b置信度 ${(r.prob * 100).toFixed(2)}%; } /script /body /htmlacceptimage/*让文件选择框只显示图片文件但这是前端提示后端仍要做校验。formData.append(image, 文件对象)的键名必须和 Flask 里request.files.get(image)的字符串一致两边拼错一个字母接口就 400。页面里同时展示置信度而不是只显示类别是为了让用户对模型判断有直观感觉答辩时可以专门点这个信息讲。要让 Flask 直接托管这个页面在app.py里加一个路由from flask import render_template app.route(/) def index(): return render_template(index.html)把index.html放到templates/目录下访问http://127.0.0.1:5000/就能看到页面。4.3 接口联调的 3 个参数与常见错误前后端联调时最容易踩的坑集中在配置参数上这里列一份检查清单检查项正确姿势报错现象最大上传大小MAX_CONTENT_LENGTH设为 8MB413 Request Entity Too Large文件格式校验检查后缀名 jpg/png前端能传、后端报错debug 模式正式联调时debugFalse请求可能执行两次debugTrue时 Flask 的自动重载会让模型被加载两遍如果加载逻辑有副作用比如打印占显存会出现诡异报错。另一个容易忽略的是图片格式用户上传的图片可能有 EXIF 旋转信息PIL 的Image.open不会自动处理必要时用ImageOps.exif_transpose(img)先转正。如果 Flask 服务和 HTML5 页面不在同一个端口fetch会触发跨域这时需要给响应加上Access-Control-Allow-Origin: *或者干脆让 Flask 同时托管页面绕开跨域。提示联调时先拿一张训练集里的图片测试确认返回的 top1 是正确类别再换测试集图片。这一步能立刻分辨是接口问题还是模型问题。5. PyQt 桌面客户端信号槽、requests 调用与打包要点PyQt 桌面端和 HTML5 端调用的是同一条 Flask 接口所以模型和推理逻辑完全不用改只换界面层。这也是三层架构的好处任何一端出问题定位范围都在界面代码里。PyQt 相比网页端的优势是文件对话框和本地文件操作更原生演示时不需要起浏览器。5.1 QMainWindow 最小骨架与信号槽连接桌面端的最小界面包含三个控件一个按钮触发选图、一个标签显示图片、一个标签显示结果。信号槽是 PyQt 的通信机制点击按钮的clicked信号连接到处理函数import sys from PyQt6.QtWidgets import (QApplication, QMainWindow, QPushButton, QLabel, QFileDialog, QVBoxLayout, QWidget) class MainWindow(QMainWindow): def __init__(self): super().__init__() self.setWindowTitle(动物图像识别) self.btn QPushButton(选择图片并识别, self) self.image_label QLabel(未选择图片, self) self.result_label QLabel(等待识别..., self) layout QVBoxLayout() layout.addWidget(self.btn) layout.addWidget(self.image_label) layout.addWidget(self.result_label) container QWidget() container.setLayout(layout) self.setCentralWidget(container) self.btn.clicked.connect(self.pick_and_predict) app QApplication(sys.argv) window MainWindow() window.show() sys.exit(app.exec())PyQt6 和 PyQt5 的差异主要是exec_()变成exec()以及部分枚举名变化写代码前先确认自己装的版本。整套界面用一个垂直布局QVBoxLayout就够了不需要设计复杂的窗口结构期末作业的评分重点在功能链路不在界面美观度。5.2 通过 requests 调用 Flask 接口并展示 Top-3pick_and_predict函数里做三件事打开文件对话框、把文件作为multipart/form-data上传、解析 JSON 结果import requests def pick_and_predict(self): path, _ QFileDialog.getOpenFileName( self, 选择图片, , 图片文件 (*.jpg *.png *.jpeg)) if not path: return with open(path, rb) as f: resp requests.post( http://127.0.0.1:5000/predict, files{image: f}, timeout10 ) data resp.json() if error in data: self.result_label.setText(错误 data[error]) return top1 data[top1] lines [f识别结果{top1[label]} 置信度 {top1[prob]:.2%}] for i, item in enumerate(data[top3], 1): lines.append(f 第{i}候选{item[label]} {item[prob]:.2%}) self.result_label.setText(\n.join(lines))requests.post的files参数会自动构造 multipart 表单键名image和 Flask 端严格对应。timeout10不能省如果 Flask 服务没启动没有超时的请求会一直挂着界面看起来像死机。这里的:.2%是 Python 的百分比格式化直接把 0.87 显示成87.00%不用手动乘 100。运行前确保 Flask 服务在 5000 端口活着桌面端和网页端可以同时连着同一个服务这也是架构设计里值得在答辩时主动提的一点。5.3 界面卡死、资源路径与打包命令桌面端最常见的两个坑一个是请求阻塞 UI一个是打包后找不到文件。Flask 推理很快本机请求通常几十毫秒返回所以简单场景直接在主线程里 requests 也没问题。但如果模型换成了 resnet50 或图片特别大请求超过 1 秒界面就会未响应。解决办法是把请求放到 QThread 里主线程只负责更新界面from PyQt6.QtCore import QThread, pyqtSignal class PredictWorker(QThread): finished pyqtSignal(dict) def __init__(self, file_path): super().__init__() self.file_path file_path def run(self): with open(self.file_path, rb) as f: resp requests.post(http://127.0.0.1:5000/predict, files{image: f}, timeout10) self.finished.emit(resp.json())把requests.post挪到run()里请求完成时通过finished信号把结果传回主线程界面就不会卡。打包用 PyInstaller 一条命令pip install pyinstaller pyinstaller --windowed --onefile main_window.py--windowed去掉控制台黑框--onefile打成单文件。打包后如果提示找不到图片或模型文件是因为 PyInstaller 解包后的临时目录和源码目录不同需要用sys._MEIPASS拼接资源路径这是桌面端打包绕不开的一个坑提前在文档说明里写上可以避免答辩现场翻车。6. 答辩前的最后一道工序Top-5 置信度与错误样本回看模型能跑通只是开始能讲清楚模型为什么对、为什么错才是加分项。最后一章介绍两个验证技巧输出 Top-5 置信度以及用混淆矩阵定位错误样本。这两件事做一遍答辩时老师问什么你都有数据回答。6.1 单张图片的 Top-5 置信度输出在 Flask 的predict里只返回了 top3做分析时可以写一个独立脚本批量跑测试集并输出每个样本的 Top-5import torch from torchvision import datasets, transforms model load_model(best_animal_resnet.pth, len(classes)) model.eval() test_ds datasets.ImageFolder(data/val, transformtransforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])) with torch.no_grad(): for path, (img, label) in zip(test_ds.imgs, test_ds): prob torch.softmax(model(img.unsqueeze(0)), dim1)[0] top5 prob.topk(5) wrong OK if top5.indices[0] label else WRONG names [f{classes[i]}:{p:.2f} for i, p in zip(top5.indices, top5.values)] print(f{wrong} gt{classes[label]} - { .join(names)})topk(5)返回的indices是类别索引values是对应概率。这个输出会直接暴露模型的犹豫点比如一张猫的图如果前五个候选里有狗且概率接近说明模型没学到猫和狗的关键区分特征需要检查训练集里这两类的图片质量或数量。6.2 混淆矩阵与错误样本定位批量收集预测结果后用 scikit-learn 一键画出混淆矩阵from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt all_labels, all_preds [], [] with torch.no_grad(): for img, label in test_ds: logits model(img.unsqueeze(0)) pred logits.argmax(1).item() all_labels.append(label) all_preds.append(pred) cm confusion_matrix(all_labels, all_preds) disp ConfusionMatrixDisplay(cm, display_labelsclasses) disp.plot(cmapBlues) plt.savefig(confusion_matrix.png, dpi150)混淆矩阵的解读重点在非对角线上的数值。比如 bird 被识别成 cat 的次数特别多就去data/val/bird里翻图片大概率是鸟的图片背景里有类似猫的物体或者类别图片数量明显不平衡。把这些错误样本连同推理概率截图放进文档说明里附上为什么错、怎么改进的分析整个作业的完成度会比只看准确率高一个档次。验证这一步做完训练、接口、双端界面和结果分析就闭环了剩下的就是整理文档说明里的运行步骤和环境配置保证换一台机器能跑起来。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门