简介:基于ResNet的动物图像分类系统是一份完整的Python期末大作业资源,面向计算机相关专业学生、深度学习者以及需要完成图像分类课程设计或毕业设计的开发者。系统将PyQt桌面客户端、Flask加HTML5网页端和PyTorch模型推理整合在一起,实现上传图片、实时识别动物类别、展示分类结果等功能,并贯通了数据集生成、模型训练、权重保存、后端接口调用和前端页面呈现的完整技术链路。资源包共27个文件,整体大小41.75MB,主要包含8个Python源文件、2个编译后的pyc文件、1个pth模型权重文件、1个HTML页面模板、11张界面效果截图以及1份README说明文档,其中train.py负责训练,predict.py执行推理,window.py搭建桌面界面,myflask.py启动Web服务,代码注释清晰、目录结构分明,便于直接运行学习和二次扩展。目前已有49人学习,对想快速搭建同类动物分类系统、理解ResNet残差网络在真实项目中的应用,或准备期末大作业答辩的读者尤为实用,能明显减少从零编码和调试的时间。
1. 一个期末大作业,为什么值得用 ResNet+Flask+PyQt 做全套
做基于 resnet 的动物图像分类系统,是 Python 期末大作业里性价比很高的一条路线。很多同学做完模型训练就停在了plt.imshow()那一步,老师看到的只是一个控制台输出,辛苦调参的痕迹完全看不出来。而这个项目把 PyTorch 训练、Flask 接口封装、HTML5 网页端和 PyQt 桌面端串在了一条完整链路上,模型有真实准确率,界面有交互操作,展示时可以从浏览器现场传图,也可以打开桌面程序点选图片,答辩的说服力完全不一样。
这套方案适合三类人:正在选 Python 课程设计题目的学生,想从“会跑通教程”进阶到“能交付一个小系统”的初学者,以及需要快速搭一个图像识别 demo 去验证想法的人。核心就一句话:ResNet 负责把图像分类这件事做对,Flask 负责把模型变成可调用的接口,PyQt 和 HTML5 负责让不懂模型的人也能直接用。
2. 模型层:ResNet 选型、数据准备与 PyTorch 训练的落地细节
2.1 ResNet18 还是 ResNet50:期末作业的选型逻辑
ResNet 家族里最常见的选择是 18 层和 50 层。动物图像分类属于粗粒度分类,猫和狗、大象和企鹅之间差异明显,不需要像区分鸟类亚种那样依赖极其细微的纹理差异。因此 ResNet18 在大多数期末场景下已经足够,训练速度快,显存占用低,CPU 也能勉强跑推理。ResNet50 的优势在于更深、特征更丰富,如果数据集中有狐狸和狼这类相似物种,50 层的上限更高,但训练时间大约翻三倍,调参不当时反而容易过拟合。
我一般这样定:数据集小于 5000 张,用 ResNet18;大于 5000 张且类别之间有相似物种,用 ResNet50。不要一开始就追深网络,期末作业的时间成本是第一位的。无论选哪个,都强烈建议用预训练权重做微调,而不是从零训练。PyTorch 里加载预训练模型就是一行调用:
import torch.nn as nn from torchvision import models # 加载在 ImageNet 上预训练过的 ResNet18 model = models.resnet18(pretrained=True) # 取出全连接层的输入维度,替换成自己的类别数 num_features = model.fc.in_features num_classes = 5 # 你是几类动物就填几 model.fc = nn.Linear(num_features, num_classes)这段代码的逻辑是:预训练模型已经把低层特征(边缘、纹理、颜色块)学好了,我们只需要把最后一层分类器换成自己的动物类别。model.fc.in_features是 ResNet18 最后一个池化层输出的特征维度,等于 512;ResNet50 则是 2048。不要硬编码这个数字,用in_features取是通用做法,换模型时不用改代码。
2.2 数据集两种凑法:手工整理文件夹还是用现成数据集
期末大作业最常见的数据集来源是 Kaggle 的 Animals-10 或者自己爬图。无论哪一种,都要整理成 torchvision 能直接读的目录结构:
data/ ├── train/ │ ├── cat/ # 存放猫的图片 │ ├── dog/ │ └── bird/ └── val/ ├── cat/ ├── dog/ └── bird/ImageFolder会自动把每个子文件夹名当作类别标签,省去手写标签映射的麻烦。用 PyTorch 的内置接口加载:
from torchvision import transforms, datasets from torch.utils.data import DataLoader train_transform = transforms.Compose([ transforms.Resize((224, 224)), # ResNet 标准输入尺寸 transforms.RandomHorizontalFlip(), # 随机翻转,增强泛化 transforms.ColorJitter(brightness=0.3, contrast=0.3), transforms.ToTensor(), # 转成张量 transforms.Normalize([0.485, 0.456, 0.406], # ImageNet 均值 [0.229, 0.224, 0.225]) # ImageNet 方差 ]) val_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) train_dataset = datasets.ImageFolder('data/train', transform=train_transform) val_dataset = datasets.ImageFolder('data/val', transform=val_transform) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)这里的Normalize参数必须和预训练权重保持一致,用 ImageNet 统计出来的均值[0.485, 0.456, 0.406]和方差[0.229, 0.224, 0.225]。这是血泪教训:很多人训练时改了均值方差,模型准确率上不去,还以为是网络的问题,其实只是数值分布不对。训练时用了RandomHorizontalFlip,验证时就不要加随机增强,所以上面单独写了val_transform。
2.3 训练脚本:微调策略与断点续训
微调有两个策略。第一是冻结前面所有层,只训练全连接层,适合数据量很小(每类几十张图)的情况;第二是全部层一起微调,学习率调小,适合数据量中等以上的情况。期末作业选第二个更稳,因为数据量通常够,效果也更好。下面是一份可以直接改的训练脚本骨架:
import torch import torch.optim as optim from torch.optim.lr_scheduler import StepLR device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-4) # 微调用小学习率 scheduler = StepLR(optimizer, step_size=5, gamma=0.1) # 每 5 个 epoch 学习率降 10 倍 EPOCHS = 15 best_acc = 0.0 for epoch in range(EPOCHS): model.train() running_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() running_loss += loss.item() * images.size(0) # 验证集上算准确率 model.eval() correct = 0 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, 1) total += labels.size(0) correct += (predicted == labels).sum().item() val_acc = 100.0 * correct / total print(f'Epoch {epoch+1}/{EPOCHS}, Loss: {running_loss/len(train_dataset):.4f}, Val Acc: {val_acc:.2f}%') # 保留最优模型,并支持断点续训 if val_acc > best_acc: best_acc = val_acc torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_acc': best_acc, 'class_to_idx': train_dataset.class_to_idx, }, 'best_model.pth')学习率 1e-4 对微调来说是经验值,太高会把预训练权重洗掉,太低收敛太慢。StepLR在 15 个 epoch 的训练里会在第 5 和第 10 个 epoch 后降低学习率,后期做更精细的参数微调。保存时把class_to_idx一起存进去,后面做推理时要知道“类别索引 0 对应的是猫还是狗”,这个映射关系只存在于训练集里,不保存的话,预测阶段就要靠猜了。
断点续训的恢复方式是把model.load_state_dict和optimizer.load_state_dict从torch.load出来的字典里取回来,同时把epoch继续往下传。期末作业虽然用不大上,但训练中途断电或者调参翻车时,这相当于后悔药,不用重头再来。
3. 服务层:用 Flask 把模型包装成可调用的 API
3.1 Flask 加载模型的时机:全局加载一次,不要在请求里反复初始化
Flask 是给 Python 期末项目做接口封装最常见的选择,轻量、无 ORM 负担、一个文件就能启动。它在这个项目里的角色是把 PyTorch 模型挡住,让 PyQt 桌面端和 HTML5 网页端都走 HTTP 请求拿到分类结果。模型加载有一个关键原则:服务启动时加载一次,存到全局变量,而不是每个请求进来都重新torch.load。一个模型文件几百 MB,每次请求都加载的话,接口延迟会从几十毫秒飙升到几秒,内存也可能被撑爆。
import torch import torch.nn as nn from torchvision import models, transforms from flask import Flask, request, jsonify from PIL import Image import io app = Flask(__name__) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 定义模型结构,与训练时保持一致 def create_model(num_classes): model = models.resnet18(pretrained=False) num_features = model.fc.in_features model.fc = nn.Linear(num_features, num_classes) return model # 启动时加载权重,只做一次 model = create_model(num_classes=5) checkpoint = torch.load('best_model.pth', map_location=device) model.load_state_dict(checkpoint['model_state_dict']) model.to(device) model.eval() # 类别映射:训练时保存的 class_to_idx 反转 idx_to_label = {v: k for k, v in checkpoint['class_to_idx'].items()}pretrained=False是关键,加载权重时不会再联网下载模型;map_location=device解决的是服务器或本机没有 GPU 时直接报错的问题,不写这一句,在 CPU 机器上torch.load会尝试把张量放到cuda:0,直接崩。model.eval()是经常被忘的一行,它会关闭 dropout 和 batch norm 的训练行为,不调用的话,同样的图每次预测结果可能都不一样。
3.2 封装预测接口:图片上传、预处理与结果返回
接口设计成一个POST /predict,接收 multipart 表单里的图片文件,返回最可能的类别和置信度。这里最容易犯的错误是拿训练时的RandomHorizontalFlip和ColorJitter直接用在推理上,导致预测结果抖动。推理时要用单独的预处理链,只做Resize、ToTensor、Normalize。
@app.route('/predict', methods=['POST']) def predict(): file = request.files.get('image') if file is None: return jsonify({'error': '缺少图片文件'}), 400 # 转成 RGB,去掉透明通道和灰度模式 img = Image.open(file.stream).convert('RGB') img = img.resize((224, 224)) # 预处理:与训练时的验证集保持一致 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) tensor = transform(img).unsqueeze(0).to(device) with torch.no_grad(): outputs = model(tensor) probs = torch.softmax(outputs, dim=1) top_prob, top_idx = torch.topk(probs, k=3) results = [{ 'label': idx_to_label[int(top_idx[0][i])], 'probability': round(float(top_prob[0][i]), 4) } for i in range(top_idx.shape[1])] return jsonify({'results': results})convert('RGB')解决了用户上传 PNG 透明图或灰度图的问题,不转的话ToTensor出来是 4 通道或 1 通道,模型直接报维度错误。unsqueeze(0)把单张图扩展成 batch 维度,PyTorch 要求输入是四维张量(batch, channel, height, width)。返回前 3 个结果而不是只返回 Top1,好处是 Web 前端可以把“最可能是猫、其次是狗”这种信息展示出来,比单纯一个标签更有说服力。置信度用softmax转成概率再保留 4 位小数,直接返回logits会让前端看到一个负数,解释成本太高。
3.3 CORS 与本地部署的坑
两个端都要访问这个 Flask 服务:PyQt 桌面端用requests库走本地127.0.0.1:5000,HTML5 页面用浏览器里的fetch直接跨域请求。浏览器对跨域有同源策略限制,从file://协议打开 HTML 页面去请求127.0.0.1:5000会被拦截,表现就是控制台报 CORS 错误,接口本身是通的。
from flask_cors import CORS CORS(app) # 允许所有来源跨域,本地开发够用flask_cors是 Flask 生态里的标准扩展,一行解决跨域。这里不展开讲 CORS 的复杂配置,期末项目放开所有来源即可。如果不想引入扩展,也可以在每次响应后手动添加Access-Control-Allow-Origin响应头,但没必要,标准扩展更省心。
启动服务时,如果只想本机访问就app.run(host='127.0.0.1', port=5000),想用手机在同一局域网里测 HTML5 页面,就改成app.run(host='0.0.0.0', port=5000)。注意debug=True虽然方便调试,但会开启 reloader,修改代码后模型会重新加载一次,对 GPU 显存来说是压力,建议开发时用,展示时关掉。
4. 展示层:PyQt 桌面端与 HTML5 网页端的双线实现
4.1 PyQt 客户端:文件选择、图片预览与结果展示
PyQt5 桌面端的定位是一个本地工具,用户选一张图片,点按钮,看到预测结果。它本质上只是 Flask 接口的一个 HTTP 客户端,不需要在 PyQt 进程里加载 PyTorch 模型,这大大降低了桌面端的启动速度和内存占用。核心代码集中在三块:选择文件、显示图片、发起请求。
import requests from PyQt5.QtWidgets import (QApplication, QWidget, QLabel, QPushButton, QFileDialog, QVBoxLayout) from PyQt5.QtGui import QPixmap class AnimalClassifierWindow(QWidget): def __init__(self): super().__init__() self.setWindowTitle('动物图像分类系统') self.setGeometry(200, 200, 600, 500) self.image_label = QLabel('请选择图片') self.image_label.setFixedSize(400, 300) self.result_label = QLabel('分类结果:') self.btn = QPushButton('选择图片并预测') self.btn.clicked.connect(self.choose_and_predict) layout = QVBoxLayout() layout.addWidget(self.image_label) layout.addWidget(self.result_label) layout.addWidget(self.btn) self.setLayout(layout) def choose_and_predict(self): path, _ = QFileDialog.getOpenFileName( self, '选择图片', '', '图片文件 (*.png *.jpg *.jpeg *.bmp)') if not path: return pixmap = QPixmap(path) self.image_label.setPixmap( pixmap.scaled(400, 300, aspectRatioMode=1)) try: with open(path, 'rb') as f: resp = requests.post( 'http://127.0.0.1:5000/predict', files={'image': f}, timeout=10 ) data = resp.json() top = data['results'][0] self.result_label.setText( f"分类结果:{top['label']},概率:{top['probability']:.2%}") except Exception as e: self.result_label.setText(f'请求失败:{str(e)}')aspectRatioMode=1对应Qt.KeepAspectRatio,缩略显示时保持图片纵横比,不会把一张长方形图片拉伸变形。requests.post里files={'image': f}的字段名必须和 Flask 端request.files.get('image')一致,一个叫image一个叫file就会拿到空值。加上timeout=10防止 Flask 没启动时客户端卡死,这是期末答辩现场最容易翻车的场景:老师点完按钮,程序转圈没反应,其实只是后端服务没起。PyQt 端在请求期间会阻塞界面,这是单线程 GUI 的通病,期末展示不深究,进阶做法是用QThread把请求放到子线程,界面就不会“卡死”。
4.2 HTML5 页面:用 FormData 调接口、展示图片与概率
HTML5 网页端是另一条展示路线,它的优势是不需要安装任何环境,浏览器打开就行。整个页面可以是一个单文件index.html,内联 CSS 和 JavaScript,便于塞进项目 zip 里直接运行。核心逻辑是:用户选择图片后,前端先本地预览,再通过fetch把图片发给 Flask,拿到结果后渲染到页面上。
<!DOCTYPE html> <html lang="zh-CN"> <head> <meta charset="UTF-8"> <title>动物图像分类系统</title> </head> <body> <h2>动物图像分类系统</h2> <input type="file" id="imageInput" accept="image/*"> <img id="preview" width="300" height="220" alt="预览区"> <div id="result"></div> <script> const input = document.getElementById('imageInput'); const preview = document.getElementById('preview'); const resultDiv = document.getElementById('result'); input.addEventListener('change', function() { const file = input.files[0]; if (!file) return; // 本地预览,不需要上传到服务器 preview.src = URL.createObjectURL(file); const formData = new FormData(); formData.append('image', file); fetch('http://127.0.0.1:5000/predict', { method: 'POST', body: formData }) .then(res => res.json()) .then(data => { if (data.error) { resultDiv.innerHTML = '<p style="color:red">' + data.error + '</p>'; return; } let html = '<ul>'; data.results.forEach(item => { html += '<li>' + item.label + ':' + (item.probability * 100).toFixed(2) + '%</li>'; }); html += '</ul>'; resultDiv.innerHTML = html; }) .catch(err => { resultDiv.innerHTML = '<p style="color:red">请求失败,请确认 Flask 已启动</p>'; }); }); </script> </body> </html>FormData构造的 multipart 表单数据,字段名image要和 Flask 端对应。URL.createObjectURL(file)是浏览器本地生成临时 URL,比把图片readAsDataURL再转字符串清爽得多,内存释放不用管,页面关闭自动回收。这里不要用axios,原生fetch已经够用,少一个依赖,期末项目少一点外部库少一点风险。如果 HTML 文件不是从 Flask 的templates目录下发的,而是直接双击打开,那么上面第 3 章的CORS(app)就必不可少,否则浏览器会拦截这个跨域请求。
4.3 把三条线串起来:一次完整的本地部署流程
一个能现场演示的系统,从上到下的启动顺序是固定的。写进项目 README 里,答辩前照做一遍:
# 第一步:启动 Flask 服务(终端窗口 1) python app.py # 看到 Serving Flask app 和 Running on http://127.0.0.1:5000 就说明服务已就绪 # 第二步:打开 HTML5 页面,直接双击或用浏览器打开 index.html # 此时浏览器里的页面可以通过 fetch 访问本机 5000 端口 # 第三步:启动 PyQt 桌面端(终端窗口 2) python gui.pyFlask 是这个系统的中枢,两个前端都只依赖它。启动顺序很重要:先跑 Flask,再开前端。如果先打开网页再启动 Flask,页面加载时没有接口可访问,但只要不点预测按钮就不会报错,等 Flask 起来了再点一样能通。PyQt 端同理。现场演示时,我一般把 Flask 窗口和 PyQt 窗口并排摆,让老师看到日志在刷请求记录,这比口头解释“我做了接口封装”直观得多。
5. 避坑指南:训练、接口与打包的五个翻车现场
5.1 推理时的预处理和训练时的验证集不一致,准确率断崖下跌
现象:训练集准确率 90% 以上,模型保存后单独测一张图,结果和训练时表现完全对不上,甚至同类图片预测出完全不同标签。
原因:训练管线里带了RandomHorizontalFlip和ColorJitter,而推理时直接复用了同一套 transform,或者反过来,推理时忘了加Normalize。数据增强和归一化是两个层面的操作,前者只在训练时用,后者每时每刻都要用。归一化缺失会让输入像素分布从 0 到 1 的区间直接进网络,预训练模型的数值分布假设被打破。
解决:把推理预处理单独写成一个函数或常量,固定为Resize((224, 224))+ToTensor()+Normalize三件套,不掺任何随机增强。前端两个入口(PyQt 和 HTML5)都调用这同一套逻辑,不要各自写一份。检查办法是打印一张输入图片预处理后的像素均值和方差,应接近 0 和 1。
5.2 没有 GPU 的机器上torch.load直接报错
现象:用 GPU 训练完,把代码和模型拷到另一台只有 CPU 的电脑上,启动 Flask 时抛RuntimeError: Attempting to deserialize object on a CUDA device。
原因:torch.save默认把张量所在设备信息写进了模型文件,GPU 上保存的权重目标设备是cuda:0,CPU 机器加载时照本宣科。
解决:加载时固定写torch.load('best_model.pth', map_location='cpu'),或在 Flask 里用map_location=device并让device = 'cuda' if torch.cuda.is_available() else 'cpu'。这是一个必须在写代码时就养成的习惯,不要指望演示那台机器一定配置正确。
5.3 Flask 的debug=True导致模型加载两次,显存溢出
现象:启动 Flask 后日志提示Restarting with stat,然后 GPU 显存占用翻倍,甚至直接CUDA out of memory。
原因:debug=True会启动 Werkzeug 的 reloader,它会重新启动一个子进程来监听文件变化。模型定义在模块顶层,父进程和子进程各加载一次,两块显存。资源吃紧时直接溢出。
解决:展示和部署时强制debug=False。开发时如果确实需要热重载,把模型加载代码挪进函数,用lru_cache或模块级全局变量保证只加载一次,但期末项目不需要这么复杂,关掉 debug 即可。
5.4 PyQt 打包成 exe 后找不到模型文件
现象:源码里python gui.py一切正常,用 PyInstaller 打成 exe 后,点击预测提示文件不存在,或者 Flask 报模型加载路径错误。
原因:PyInstaller 打包后,工作目录不是 exe 所在目录。模型文件没有被识别为资源打进包里,或者代码用了相对路径'best_model.pth',运行时当前目录和源码目录不是同一个。
解决:两种方案任选。一是打包时把模型作为外部文件放在 exe 旁边,代码里用绝对路径拼接:
import sys import os def get_model_path(): if getattr(sys, 'frozen', False): # 打包成 exe 后,以 exe 所在目录为基准 base_dir = os.path.dirname(sys.executable) else: base_dir = os.path.dirname(os.path.abspath(__file__)) return os.path.join(base_dir, 'best_model.pth')二是用 PyInstaller 的--add-data把模型打进包内,运行时通过sys._MEIPASS临时目录访问。期末答辩推荐第一种,模型文件外置,换模型不用重新打包,出了问题也好排查。
5.5 端口被占用导致 Flask 启动失败
现象:启动 Flask 时日志报Address already in use,或者页面请求一直连不上127.0.0.1:5000。
原因:上一次运行的服务没有正常退出,或者有其他程序占用了 5000 端口。Windows 上常见于之前Ctrl+C没关干净,后台进程还在监听。
解决:换端口启动最简单,app.run(port=5001);或者查到占用进程并结束它。Windows 查占用:
netstat -ano | findstr :5000 taskkill /PID <pid> /F如果平时电脑上跑过其他 Flask 项目,5000 被占是常态。因此我建议 Flask 启动时自动寻找可用端口,或者在 README 里明确写出端口冲突的处理命令,答辩时处理起来不手忙脚乱。
6. 进阶:想涨准确率、想拿高分,可以做的三件小事
如果基础链路已经跑通,想在期末答辩里获得更好评价,有低中高三个成本的动作可以做。低成本的是换一个更大的预训练模型。把models.resnet18(pretrained=True)换成models.resnet50(pretrained=True),其他代码都不用改,类别数从 5 到 50 都能适应,因为in_features是自适应取的。ResNet50 在相似动物上表现更好,缺点是训练时间和显存占用上升,如果你的电脑跑得动,这个改动性价比很高。
中等成本的是给模型加一个可视化解释模块,用 Grad-CAM 生成热力图展示模型重点关注的区域。答辩时上传一张猫的图片,界面上除了显示“猫,置信度 92%”,再展示一张热力图,猫脸位置发红,背景发蓝。这直接说明模型不是靠蒙的,而是学对了特征。PyTorch 里实现 Grad-CAM 需要注册 hook 取最后一个卷积层的梯度,代码大约 30 行,网上现成封装很多,在期末项目里属于“做了就明显超出平均水平”的一档。
高成本但也是最能体现工程能力的是把前端从“手动传图”升级成“从摄像头实时识别”。用 HTML5 的getUserMedia获取摄像头流,定期截帧发送到 Flask 推理。这个改动涉及到帧率控制和请求频率问题,要做好防抖,否则每帧都发请求后端撑不住。做出来以后,项目演示从“点选图片”变成“对着镜头里的动物直接识别”,体验质变。
说一个我自己踩过的坑:做完基础版本后盲目追求准确率,加了很多数据增强和复杂的调度策略,结果训练时间翻了几倍,答辩前夜还在调整参数,最终的准确率提升不到 2 个点。后来明白,期末作业的评分维度是多方面的,工程完整度、界面交互、部署可运行性,和准确率哪怕提升到 98%,都不如“系统能稳定跑通、前后端配合顺畅、展示过程不翻车”来得重要。先把链路跑稳,再谈涨点,顺序不要颠倒。希望我这几条经验对你有帮助,照着这个思路把项目搭起来,答辩时就有东西可讲,有底气可站。
本文还有配套的精品资源,点击获取