ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

基于Python与CNN的图像分类系统:五大模型双框架毕设实战

基于Python与CNN的图像分类系统:五大模型双框架毕设实战 简介一份面向计算机相关专业学生与初学者的CNN图像分类系统完整实现覆盖毕业设计、课程设计、项目演示等典型场景。压缩包共22个文件包含13个Python脚本、2个pyc编译文件、打包好的数据集与训练模型、说明文档及Web前端配套文件总大小仅62KB便于快速下载与二次开发。项目内置LeNet-5、AlexNet、GoogLeNet、ResNet等经典卷积神经网络并分别给出TensorFlow与PyTorch两种框架实现同时配有图像分类的Web交互界面、模型定义、评估与矩阵分析脚本可帮助读者从数据准备、模型训练到界面部署走通完整流程。资源中的说明文档和分类索引文件让结构更清晰适合直接用于毕设/课设答辩演示也可作为深度学习入门进阶的实战参考。目前已有132人浏览学习代码经测试运行通过项目曾获导师认可且答辩评分95分。1. 基于Python与CNN的图像分类毕设这套资源解决什么问题做过毕设的人都知道图像分类最容易卡在「模型跑得动但精度上不去」「代码跑通了却讲不明白原理」这两道坎上。这套基于 Python 的 CNN 图像分类系统不是零碎示例而是把 LeNet-5、AlexNet、GoogLeNet、ResNet 五种经典卷积神经网络同时封装好并分别给出 TensorFlow 和 PyTorch 两套框架版本还附带训练好的模型权重、数据集压缩包、class_indices.json 标签映射和一套 Flask Web 交互界面。拿到手之后你可以直接启动网页上传图片做分类也可以翻源码、改参数、重新训练再配合说明文档去写论文和准备答辩。它适合软件工程、计科、人工智能、自动化、电子信息等专业的学生做毕业设计或课程设计也适合想尽快把 CNN 全流程跑通再看细节的入门者。2. 五种经典CNN模型逐个拆LeNet-5到ResNet的选型逻辑与参数量差异2.1 LeNet-5与AlexNet经典结构为什么到现在还在用LeNet-5 是 1998 年由 Yann LeCun 提出的最初用于手写数字识别输入是 32×32 的灰度图。它的意义在于把卷积、池化、全连接三个组件拼成了一个完整范式后续所有 CNN 都是在这个骨架上扩展出来的。在毕设论文里LeNet-5 非常适合放在「从原理到实现」的第一章用它讲清卷积核怎么滑动、特征图尺寸怎么算、池化为什么能降采样这些原理问题是答辩时老师最常追问的。资源里把它放在编号 01 的位置我理解它的作用是当基线模型用。先跑通 LeNet-5确认数据加载、训练、保存、加载这条链路是通的再换成更深的网络。LeNet-5 参数量只有约 6 万CPU 上几分钟就能收敛非常适合用来排查环境问题。如果你的数据集不大比如每类只有几百张图用 LeNet-5 做主力模型其实也说得过去关键是你能讲清楚为什么选它而不是无脑上 ResNet。AlexNet 是 2012 年 ImageNet 竞赛的冠军模型它真正让 CNN 火了起来。它的贡献不止是深度还包括 ReLU 激活函数、Dropout 随机失活、GPU 并行训练和数据增强。从参数量上看AlexNet 约有 6000 万参数比 LeNet-5 大了两个数量级所以训练时对显存和数据量的要求也明显更高。在毕设场景里AlexNet 的价值在于它展示了「怎么把网络做大还能正常收敛」Dropout 和数据增强这两招在后来的实战中几乎必用。2.2 GoogLeNet与ResNet深度之外的另一条路GoogLeNet 出现在 2014 年它的核心是 Inception 模块。一个 Inception 模块里同时用 1×1、3×3、5×5 三个尺寸的卷积核并行提取特征再把结果沿通道方向拼接。这样做的好处是网络可以变得更「宽」同一个层能捕捉不同尺度的模式。为了控制计算量模块里先用 1×1 卷积做通道降维这是非常经典的操作。资源中把它放在编号 04我的理解是它用于在 LeNet-5、AlexNet 之后进一步提升精度属于中间选择。ResNet 则是 2015 年的模型核心是残差连接把输入直接加到卷积层输出上即输出 F(x) x。这个操作看似简单却解决了深度网络退化问题让网络可以堆到 50 层以上。ResNet 系列里 ResNet-50 是最常用的版本参数量约 2500 万在 ImageNet 上的表现远好于同期模型。实际使用中如果你的毕设没有特殊的结构创新要求我的建议是直接以 ResNet 为主力模型。它训练稳定、精度高、答辩时也容易讲——「我用残差解决了梯度消失问题」这句话本身就是得分点。GoogLeNet 则可以作为对比实验出现用来展示你做了多模型横向比较这比单跑一个模型的论文丰富得多。下表把这五种模型的参数规模和适用场景做了个汇总方便你写论文时直接引用。模型提出年份核心创新典型参数量适合扮演的角色LeNet-51998卷积池化全连接范式约 6 万基线模型、原理讲解AlexNet2012ReLU、Dropout、双GPU约 6000 万经典对比模型GoogLeNet2014Inception模块、1×1降维约 700 万对比实验、效率分析ResNet2015残差连接、BNResNet-50约2500万主力模型、精度担当2.3 双框架模型定义对照PyTorch与TensorFlow怎么选这份资源给我印象最深的一点是同一个模型目录下同时有 PyTorch 和 TensorFlow 两套实现。PyTorch 的代码更接近 Python 直觉定义模型就像写类一样自然适合快速改结构做实验TensorFlow 则更适合工业部署SavedModel 格式在服务端跑起来非常方便。毕设阶段我一般建议主用 PyTorch原因有两点一是调试时打印中间张量尺寸很方便二是绝大多数的论文复现代码都默认 PyTorch遇到问题搜解决方案更容易。下面是资源里 PyTorch 版 AlexNet 的完整定义节选注意看特征提取层和分类层的分工# model.py 中 PyTorch 风格的 AlexNet 定义节选 import torch.nn as nn class AlexNet(nn.Module): def __init__(self, num_classes5): super(AlexNet, self).__init__() self.features nn.Sequential( nn.Conv2d(3, 96, 11, stride4, padding2), # 224 - 55 nn.ReLU(inplaceTrue), nn.MaxPool2d(3, 2), # 55 - 27 nn.Conv2d(96, 256, 5, padding2), # 27 - 27 nn.ReLU(inplaceTrue), nn.MaxPool2d(3, 2), # 27 - 13 nn.Conv2d(256, 384, 3, padding1), # 13 - 13 nn.ReLU(inplaceTrue), nn.Conv2d(384, 384, 3, padding1), # 13 - 13 nn.ReLU(inplaceTrue), nn.Conv2d(384, 256, 3, padding1), # 13 - 13 nn.ReLU(inplaceTrue), nn.MaxPool2d(3, 2), # 13 - 6 ) self.classifier nn.Sequential( nn.Dropout(), nn.Linear(256 * 6 * 6, 4096), nn.ReLU(inplaceTrue), nn.Dropout(), nn.Linear(4096, 4096), nn.ReLU(inplaceTrue), nn.Linear(4096, num_classes), ) def forward(self, x): x self.features(x) x x.view(x.size(0), -1) x self.classifier(x) return x这里的关键点是features 部分负责把空间信息逐步抽象成高维特征最后通过池化得到 256 通道的 6×6 特征图classifier 部分把特征图展平后接三个全连接层最后输出 num_classes 个类别的得分。注释里的 224 - 55 - 27 表示特征图尺寸随卷积和池化的变化改输入尺寸时一定要重新算一遍否则 view 那行会报维度错误。TensorFlow 版的 LeNet-5 定义则要更精简Sequential 方式一目了然# model.py 中 TensorFlow 风格的 LeNet-5 定义 import tensorflow as tf from tensorflow.keras import layers, models def lenet5(input_shape(32, 32, 3), num_classes5): model models.Sequential([ layers.Conv2D(6, 5, activationrelu, input_shapeinput_shape), layers.MaxPool2D(pool_size2, strides2), layers.Conv2D(16, 5, activationrelu), layers.MaxPool2D(pool_size2, strides2), layers.Flatten(), layers.Dense(120, activationrelu), layers.Dense(84, activationrelu), layers.Dense(num_classes, activationsoftmax) ]) return model对比两套代码可以看出同样的结构PyTorch 需要显式写 forward 传播TensorFlow 用 Sequential 就能拼完。我的血泪经验是不要在一开始就纠结框架选择先选定一个框架把数据链路跑通另一个框架作为「论文里写一嘴对比」的存在就够了。3. 数据准备与训练脚本从目录结构到收敛判断的落地参数3.1 目录结构、数据加载与标签映射拿到的压缩包解压后数据集部分要按训练、验证、测试三块组织起来这是图像分类项目最常见也是最少出错的组织方式。目录结构如下dataset/ ├── train/ │ ├── class_cat/ │ │ ├── cat_001.jpg │ │ ├── cat_002.jpg │ │ └── ... │ ├── class_dog/ │ │ ├── dog_001.jpg │ │ └── ... │ └── ... ├── val/ │ ├── class_cat/ │ └── ... └── test/ └── ...PyTorch 的 torchvision.datasets.ImageFolder 可以直接读取这种结构它会按子目录名自动生成类别索引。资源里的 class_indices.json 就是这份索引的映射文件内容形如{ class_cat: 0, class_dog: 1, class_bird: 2, class_fish: 3, class_other: 4 }这个文件有两个作用训练时告诉 DataLoader 类别顺序部署时把预测的整数索引反查回真实类别名。我的习惯是训练开始前就把 class_indices.json 打印出来和图片目录比对一遍防止目录名里有隐藏字符或重复类别这一步能避免后面百分之八十的错乱问题。下面是数据加载和预处理的完整写法# 数据加载与预处理PyTorch 版本 import json import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms # 读取类别映射用于训练后对照 with open(class_indices.json, r, encodingutf-8) as f: class_indices json.load(f) num_classes len(class_indices) # 训练集加数据增强 train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 验证/测试集只做缩放和归一化 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 datasets.ImageFolder(dataset/train, transformtrain_transform) val_dataset datasets.ImageFolder(dataset/val, transformval_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4)这里有个参数需要重点理解Normalize 的 mean 和 std 用的是 ImageNet 数据集的统计值。如果你的数据集是自采的自然图像直接用这套值问题不大但如果是医学影像、遥感图这类分布差异大的数据建议用自己数据集的均值和标准差否则第一层卷积的输入分布会有偏移表现为训练初期 loss 很难降。3.2 训练脚本关键配置batch_size、学习率与 epoch 怎么设训练脚本是整个工程的心脏。对于这个资源里的五类分类任务我一般会给出这样的初始配置参数建议值调参方向batch_size32显存不足就降到 16 或 8学习率Adam0.001不收敛就降到 0.0001num_epochs30观察 val_loss 是否还有下降趋势优化器Adam换 SGD momentum 需配合 lr 调整损失函数CrossEntropyLoss多分类标配训练循环的核心代码如下# 训练循环核心片段PyTorch device torch.device(cuda if torch.cuda.is_available() else cpu) model AlexNet(num_classesnum_classes).to(device) criterion torch.nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) best_val_acc 0.0 for epoch in range(30): model.train() running_loss 0.0 correct 0 total 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() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() train_acc 100 * correct / total print(fEpoch [{epoch1}/30], Loss: {running_loss/len(train_loader):.4f}, fTrain Acc: {train_acc:.2f}%) # 每个 epoch 后验证一次保留最佳模型 val_acc evaluate(model, val_loader, device) if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), models/best_model.pth) print(f模型已保存当前最佳验证准确率 {best_val_acc:.2f}%)代码里没有直接写 evaluate 函数的定义它在资源里通常位于 model.py 或单独的工具模块中核心逻辑和训练循环里的前向计算一致只是不反向传播。这里最值得留心的一行是torch.save(model.state_dict(), ...)它只保存了模型参数状态没有保存模型结构。后面每次使用都要先实例化一个相同结构的模型再加载参数这个细节在第五章会专门展开。3.3 怎么判断训练是否正常Loss曲线与准确率并不是同一件事很多同学看到训练准确率一路上升就觉得没问题其实这是最常见的错觉。准确率对参数变化不敏感尤其在训练后期acc 可能从 90% 慢慢爬到 92%但 loss 的波动能更早暴露问题。我训练时会同时盯着几个信号第一训练 loss 是否在前 5 个 epoch 内快速下降。如果 loss 几乎不动先检查学习率是不是太小或者数据加载的标签是否错位。第二训练准确率和验证准确率的差值。如果两者差距超过 10 个百分点基本可以断定过拟合需要增强数据增强强度或加大 Dropout。第三val_loss 是否在某个 epoch 后开始回升而 val_acc 还在涨这说明模型开始过拟合了最佳模型应该在 val_loss 最低点保存而不是在最后一个 epoch 保存。资源里附带的说明文档如果有一张 loss 曲线图那通常是从 TensorBoard 或 matplotlib 生成的。我建议你复现训练时把它一起做了答辩 PPT 上摆一张平滑的 loss 收敛曲线比任何文字描述都有说服力。4. Flask Web端部署把训练好的模型封装成网页分类服务4.1 main.py 的路由设计与推理流程这个资源的加分项之一是带了一个可直接运行的 Flask 应用。它不是那种「跑通就完事」的命令行工具而是有一个网页界面能上传图片、看到分类结果和置信度。这在毕设答辩现场演示时效果非常好。main.py 的核心路由设计如下# main.py 核心接口Flask 路由与推理 import io import json import torch from flask import Flask, request, jsonify, render_template from PIL import Image from model import ResNet import torchvision.transforms as transforms app Flask(__name__) # 全局加载模型避免每次请求重复加载 device torch.device(cuda if torch.cuda.is_available() else cpu) model ResNet(num_classes5) model.load_state_dict(torch.load(models/resnet.pth, map_locationdevice)) model.to(device) model.eval() # 加载类别映射 with open(class_indices.json, r, encodingutf-8) as f: class_indices json.load(f) class_names list(class_indices.keys()) # 预处理必须与训练时一致 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]) ]) app.route(/) def index(): return render_template(index.html) app.route(/predict, methods[POST]) def predict(): file request.files.get(image) if file is None: return jsonify({error: 未上传图片}), 400 img Image.open(io.BytesIO(file.read())).convert(RGB) input_tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) probs torch.softmax(output, dim1) prob, idx torch.max(probs, dim1) result { label: class_names[idx.item()], probability: round(prob.item(), 4) } return jsonify(result) if __name__ __main__: app.run(host0.0.0.0, port5000, debugFalse)这里有两个值得注意的设计。一是模型在全局加载且只加载一次代码里没有把它写进 predict 函数内部否则每次请求都要重新读权重文件推理耗时会被拉长到秒级。二是推理被包在torch.no_grad()里不追踪梯度推理速度更快、内存占用更小。unsqueeze(0)的作用是把单张图片的 [3, 224, 224] 张量变成 [1, 3, 224, 224]模拟出一个 batch 大小为 1 的输入因为模型要求输入必须带 batch 维。4.2 预处理必须和训练保持一致Resize、Normalize、通道顺序部署环节最隐蔽的坑就是预处理不一致这是血泪经验。训练时用的是 CenterCrop 224部署时如果随手写成直接 Resize 到 224虽然图片尺寸相同但内容分布有差异推理结果可能偏差好几个百分点。我见过有人训练准确率 90%部署后准确率掉到 70%最后发现是 Resize 和 CenterCrop 的差异造成的。同样要认真核对的是 Normalize 参数和通道顺序。PyTorch 读入的 PIL 图像会通过 ToTensor 变成 CHW 顺序 [3, 224, 224]但如果用了 OpenCV 的 imread读出来是 HWC 且通道顺序是 BGR必须先转换成 RGB 再做变换。资源里静态文件如果有一段单独的预处理工具函数你直接用就好如果是自己改一定要把下面的检查项逐条对一遍检查项训练时代码部署时代码不一致后果缩放方式Resize(256)CenterCrop(224)同样特征分布异常归一化均值0.485, 0.456, 0.406同样首层输入偏移归一化方差0.229, 0.224, 0.225同样首层输入偏移通道顺序RGBRGB颜色错乱4.3 Matrix.py 混淆矩阵评估报告怎么生成Matrix.py 这个文件名字看起来不起眼实际是答辩时的重要素材。它负责在测试集上统计每个类别的预测情况生成混淆矩阵并辅助计算精确率、召回率、F1 值。下面是一个典型的实现# Matrix.py 混淆矩阵计算与可视化 import numpy as np import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix def plot_confusion_matrix(y_true, y_pred, class_names, save_pathconfusion_matrix.png): cm confusion_matrix(y_true, y_pred) plt.figure(figsize(8, 6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclass_names, yticklabelsclass_names) plt.xlabel(Predicted Label) plt.ylabel(True Label) plt.tight_layout() plt.savefig(save_path, dpi300) plt.show() # 输出每个类别的精确率、召回率、F1 from sklearn.metrics import classification_report def print_report(y_true, y_pred, class_names): report classification_report( y_true, y_pred, target_namesclass_names, digits4) print(report) return report调用这段代码前需要先跑一遍测试集把所有样本的真实标签和预测标签分别收集成两个列表然后把两个列表传进去。论坛上很多人抱怨混淆矩阵画不出来十有八九是没有先做预测这一步就直接传了原始数据。从答辩的角度看混淆矩阵能回答一个关键问题「你的模型在哪些类别之间最容易混淆」如果猫和狗互相误判说明特征学习还不够细致如果某一类被大面积误判为另一类可能是数据不平衡。这些分析写进论文里比单纯报一个准确率数字要高级得多。5. 避坑手册CNN毕设里最常见的五个翻车现场与排查路径5.1 显存不足OOM与 batch_size 的玄学现象训练刚开始几个 stepPyTorch 直接报CUDA out of memory或者 TensorFlow 报显存耗尽错误电脑直接卡死。原因默认 batch_size32 加上 224×224 的输入在 4GB 显存的笔记本 GPU 上很容易撑不住。很多同学以为是模型太大其实问题出在 batch 和 num_workers 的组合上num_workers 开大了内存也可能爆。解决最直接的办法是把 batch_size 从 32 降到 16 或 8显存占用几乎是线性下降。其次是检查是不是同时开了多个程序占用显存用nvidia-smi看一眼。还有一个隐藏因素验证阶段有时候忘了写torch.no_grad()导致反向传播图被保留显存被撑爆。# 查看 GPU 显存占用确认是否是其他进程抢占 nvidia-smi5.2 训练集准确率高但验证集明显偏低过拟合的典型特征现象训练集准确率一路涨到 95%验证集却只有 60%而且 val_loss 在第 15 个 epoch 后开始回升。原因模型参数量大而数据量小学到的更多是训练集的细节噪音而不是通用特征。在 LeNet-5、AlexNet 这类模型上加上少量数据增强过拟合几乎是必然的。解决先检查数据集每类图片数量如果每类少于 500 张优先考虑用预训练权重做迁移学习而不是从头训练。其次把 RandomResizedCrop、RandomHorizontalFlip、ColorJitter 这些增强手段全部打开。再不行就把 Dropout 概率从 0.5 提到 0.6并在全连接层之前加一层 Global Average Pooling 减少参数量。这些都是毕设项目里改动成本低、见效快的手段。5.3 模型文件加载报错state_dict 结构不匹配与路径问题现象load_state_dict 时一连串报错提示Missing key(s)或Unexpected key(s)也有直接提示文件不存在的。原因大部分情况是模型结构不一致。比如训练时保存的是 ResNet50 但有 5 个类别部署时实例化模型时 num_classes 传的是 10最后一层全连接维度对不上。另一个常见原因是torch.save(model.state_dict())和torch.save(model)混用前者是权重字典后者是完整模型对象。解决实例化模型时务必确保 num_classes 和训练时一致如果是权重字典就统一用load_state_dict如果是完整模型就统一用torch.load。路径方面建议在 main.py 里用绝对路径或基于os.path.dirname(__file__)拼接路径不要依赖终端当前目录否则在 database 里更换启动目录就会翻车。# 推荐基于当前文件路径拼接模型路径 import os BASE_DIR os.path.dirname(os.path.abspath(__file__)) model_path os.path.join(BASE_DIR, models, best_model.pth)5.4 Flask 部署后首次推理非常慢模型加载与 CPU 推理的耗时现象网页启动很快但点击上传图片后要等好几秒才出结果第二次就快一些。原因第一次推理耗时包含了模型权重的磁盘读取、模型参数初始化、输入预处理和真正的前向计算。如果服务器只有 CPUResNet50 单张 224×224 的推理时间就可能超过 1 秒。还有一个隐藏因素模型加载时没有设置成 eval 模式Dropout 层在推理时仍然随机失活导致结果不稳定。解决在模型加载完成后立即调用model.eval()。另外可以做一个预热在应用启动后主动跑一次假推理把模型参数和内存页先激活。我在项目里习惯加一段预热代码# 启动时预热用一个全零张量跑一次前向避免首次请求过慢 with torch.no_grad(): dummy torch.zeros(1, 3, 224, 224).to(device) _ model(dummy)如果数据集本身不太大还可以把 torch 退回到 CPU 推理并开启torch.set_num_threads(4)在精度几乎不变的前提下把速度提上去。5.5 图片读取失败但文件明明存在中文路径与编码的坑现象训练时一部分图片报FileNotFoundError或 PIL 直接报无法打开而且报错的文件名里带中文但用资源管理器能看到文件确实存在。原因Windows 下 PIL 的Image.open在读取含中文路径的文件时底层用的是系统默认编码处理 UTF-8 路径容易出问题。部分压缩包解压后目录名带中文加上文件名也是中文概率就更高了。此外还有 EXIF 方向信息导致图片旋转的问题虽然不影响读取但会影响训练效果。解决把数据集路径和文件名统一改成英文是治本的办法。如果不想改名可以在读取端用 OpenCV 代替 PILOpenCV 对中文路径的兼容性稍好一些但它读出来是 BGR 通道要注意转换。数据加载是毕设项目里最不该花时间的环节建议一开始就用英文命名。6. 迁移学习微调与模型验证把准确率再往上走一截的实操如果你用的是自建的数据集每类图片量在几百张这个量级从头训练的准确率通常不会太理想。这时候迁移学习是性价比最高的优化手段加载在 ImageNet 上预训练好的 ResNet 权重然后把最后一层全连接替换成你自己的类别数冻结前面所有层只训练分类头。这样即使只有两百张训练图片准确率也能稳定在 80% 以上。# 迁移学习微调加载 ImageNet 预训练权重替换最后一层 from torchvision import models # 加载预训练 ResNet50 model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1) # 替换最后一层全连接5 是实际类别数 num_features model.fc.in_features model.fc torch.nn.Linear(num_features, 5) # 冻结特征提取层只训练新加的分类头 for param in model.parameters(): param.requires_grad False for param in model.fc.parameters(): param.requires_grad True上面的代码关键是最后两段循环先把所有参数requires_grad设为 False再把新加的全连接层参数设回 True。训练时优化器只传入model.fc.parameters()就不会动前面那些预训练好的卷积层了。训练 10 到 15 个 epoch学习率用 0.0001通常就能超过从头训练 30 个 epoch 的效果。注意weightsmodels.ResNet50_Weights.IMAGENET1K_V1是当前稳定可用的官方权重入口老版本代码里写的pretrainedTrue在新版 torchvision 里已经不建议使用。微调之后一定要回到第五章提到的验证流程先跑测试集收集所有真实标签和预测标签传入 Matrix.py 的plot_confusion_matrix函数看混淆矩阵的热力图上有没有明显的横条纹。如果某一行的样本被大量分到其他类检查这类图片是不是本身就和别的类视觉上很像比如背景相似的猫和狗或者不同拍摄角度的同一物体。我自己的习惯是每次调完参数强制把 train、val、test 三个阶段全部重跑一遍确认数据没有泄漏、预处理完全一致、混淆矩阵和准确率都能对上之后才敢把数字写进论文。这个流程看着繁琐但每次答辩前它都救了我一次。希望帮到你。本文还有配套的精品资源点击获取
返回列表