尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

从零训练CNN水果分类模型:Python与PyTorch图像分类实践

从零训练CNN水果分类模型:Python与PyTorch图像分类实践 简介卷积神经网络CNN是深度学习中最经典的图像识别模型通过层级卷积与池化自动提取图像中的边缘、纹理等特征实现高效分类。在图像分类任务中结合迁移学习和数据增强可以有效缓解小数据集带来的过拟合问题。以水果图片分类为例从数据清洗、尺寸归一化到模型选型与训练调优完整梳理了一套可落地的图像分类工程实践。无论你是初学者还是工程师都能从中获得可复现的PyTorch实现思路。 很多人第一次用Python做图像识别上来就想着拿一个预训练模型跑通或者直接调API识别几张图觉得这样就算会了。但真正自己从零搭建一个CNN处理图片数据集把训练跑完再调参优化才算是把这套东西吃透了。这篇文章我完整记录一次用Python和卷积神经网络训练水果图片分类模型的过程包含数据集处理、网络结构设计、训练细节、踩坑记录以及最后的实际效果评估。整个项目我是在自己电脑上跑通的硬件条件不算好就是一块普通的消费级显卡所以整个流程对绝大多数想入门深度学习图像分类的读者都有参考价值。先说项目背景目标是训练一个模型能够识别常见水果的类别比如苹果、香蕉、葡萄、橙子、猕猴桃等。数据来源是公开的水果图片数据集里面按类别分好了文件夹每张图片尺寸不统一光照、角度、背景都有差异这正好是真实场景下最常见的情况也是初学者最容易忽略的问题——很多人以为模型效果差是网络结构不够深但其实大部分问题出在数据本身没有处理好。整个项目做完我的感受是CNN图像分类的入门门槛真不高但要把准确率做到可用水平考验的更多是工程细节。下面我把每一步怎么做的、为什么这样做、踩了哪些坑全部拆开来讲清楚。1. 环境准备与数据集初探1.1 本机环境与依赖清单先说环境。我用的Python版本是3.9深度学习框架选择了PyTorch。很多新手会纠结到底选TensorFlow还是PyTorch我的建议很简单如果你不是公司有统一技术栈要求自己学习就直接上PyTorch。原因是PyTorch的调试体验更友好报错信息更可读而且现在学术界和新出的预训练模型基本都是PyTorch版本遇到问题搜解决方案也更容易。我的环境配置如下操作系统Windows 11Python版本3.9.18CUDA版本11.8PyTorch版本2.0.1cu118显卡NVIDIA GeForce RTX 3060 12GB内存16GB安装依赖的方式很简单直接用pip安装pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install numpy pandas matplotlib scikit-learn tqdm pillow这里有个细节值得单独说说就是PyTorch版本和CUDA版本的匹配问题。很多人装完PyTorch跑起来发现用不了GPU多半就是CUDA版本对不上。你在NVIDIA官网查一下自己显卡支持的CUDA版本然后去PyTorch官网选择合适的安装命令。RTX 3060支持CUDA 11.8到12.x都没问题但如果你用的是老一点的显卡可能就得选低版本的CUDA。怎么确认PyTorch能正确调用GPU跑一下这段代码import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))输出True和你的显卡型号就说明GPU环境没问题。这一步很重要因为后面训练CNN如果只能用CPU训练速度会慢到你怀疑人生。1.2 水果数据集结构与特征分析我用的是网上公开的Fruits Recognition数据集里面包含了几十种水果的彩色图片我从中挑选了5类来做分类任务苹果、香蕉、葡萄、橙子、猕猴桃。选择这5类是因为它们在形状、颜色上有区分度既有挑战性又不会难到让新手直接放弃。先看数据集的目录结构这是很多教程不会细讲但实际非常重要的部分fruits/ ├── apple/ │ ├── 1.jpg │ ├── 2.jpg │ └── ... ├── banana/ │ ├── 1.jpg │ └── ... ├── grape/ │ ├── 1.jpg │ └── ... ├── orange/ │ ├── 1.jpg │ └── ... └── kiwi/ ├── 1.jpg └── ...每类水果大约有200-300张图片总共1000多张。图片的分辨率从几百像素到上千像素不等有的是正方形有的是长方形有的图片里水果占据画面主体有的图片里水果只占画面的一部分背景也比较杂乱。拿到数据集以后我做的第一件事不是直接开写代码而是先把每类图片的数量、尺寸分布统计出来同时随机抽样看一批图片。这个习惯很重要因为数据集不平衡或者图片质量问题会直接影响后面的训练效果。我统计了一下苹果234张、香蕉256张、葡萄221张、橙子245张、猕猴桃198张各类别数量基本接近没有严重的类别不平衡问题。但图片尺寸差异非常大最小的只有约300x300最大的超过2000x2000这决定了我们必须对所有图片做统一的预处理。另外我还发现了几个数据集中常见的脏数据问题有些图片其实是同一张图的不同缩放版本有些图片带水印或Logo个别图片甚至根本不含水果比如有一张苹果文件夹里的图拍的是一块石头应该是标注时候不小心混进去的。这些问题如果不去处理模型就会在训练时接收到错误信息导致准确率下降。2. 数据预处理这一步比模型设计更影响效果2.1 图片清洗与格式统一很多新手直接装好库就开始训练结果死活跑不出理想的准确率然后怀疑是网络结构有问题。其实大部分时候问题出在数据没有处理好。我的做法是先写一个脚本对图片做基础清洗。第一步是检查每张图片是否能被PIL正常打开打不开的直接删除或单独放一边。第二步是查看每张图片的通道数确保都是RGB三通道。第三步是去重用文件的MD5哈希值找出完全重复的图片。import os from PIL import Image import hashlib def clean_dataset(root_dir): bad_images [] duplicate_images [] seen_hashes {} for class_name in os.listdir(root_dir): class_path os.path.join(root_dir, class_name) if not os.path.isdir(class_path): continue for img_name in os.listdir(class_path): img_path os.path.join(class_path, img_name) try: with Image.open(img_path) as img: img.verify() # 验证图片是否完整 # 计算MD5去重 with open(img_path, rb) as f: file_hash hashlib.md5(f.read()).hexdigest() if file_hash in seen_hashes: duplicate_images.append(img_path) else: seen_hashes[file_hash] img_path except Exception as e: bad_images.append(img_path) return bad_images, duplicate_images跑完这个脚本我删除了3张损坏图片和12张重复图片剩下的图片基本都是可用的。这个步骤虽然看起来不直接贡献模型效果但实际意义非常大因为一张损坏的图片可能在训练过程中直接触发异常导致整个训练中断或者在数据加载时报错浪费你大量排查问题的时间。2.2 尺寸归一化与数据增强策略CNN对输入图片的尺寸有固定要求常见的有224x224、256x256等。我选的ResNet系列网络默认输入是224x224所以要把所有图片都缩放到这个尺寸。但这里有一个细节就是直接resize会改变图片的宽高比导致水果形状变形。比如一张长方形的苹果照片直接强制压缩成224x224苹果就会变成扁的。虽然CNN理论上对这种变形有一定容忍度但更好的做法是先把图片的短边缩放到224再做中心裁剪到224x224。这样既统一了尺寸又最大程度保留了水果的主体特征。数据增强也是关键一环。我见过不少新手完全不做数据增强结果模型在训练集上准确率很高一到测试集就掉链子这就是过拟合。数据增强的本质是无中生有地扩充一个模型的训练数据量同时提升模型的泛化能力。我使用的数据增强策略如下随机水平翻转概率0.5随机旋转-15度到15度随机仿射变换scale 0.8到1.2随机色彩抖动亮度、对比度、饱和度调整归一化到[0, 1]区间需要留意的是归一化时使用的均值和标准差最好是基于自己数据集计算出来的而不是直接套用ImageNet的默认值。虽然ImageNet的均值标准差在很多任务上也能用但如果你的图片整体风格差异很大用自己的数据算出来的值效果会更好。计算自己数据集均值和标准差的代码import numpy as np from PIL import Image import os from tqdm import tqdm def compute_mean_std(root_dir): means np.zeros(3) stds np.zeros(3) total_count 0 for class_name in os.listdir(root_dir): class_path os.path.join(root_dir, class_name) if not os.path.isdir(class_path): continue for img_name in tqdm(os.listdir(class_path)): img_path os.path.join(class_path, img_name) img Image.open(img_path).convert(RGB) img img.resize((224, 224)) img_array np.array(img) / 255.0 means img_array.mean(axis(0, 1)) stds img_array.std(axis(0, 1)) total_count 1 means / total_count stds / total_count return means, stds我算出来的均值和标准差大概是[0.482, 0.442, 0.363]和[0.234, 0.239, 0.242]和ImageNet的[0.485, 0.456, 0.406]确实有差异用了自己的数值后模型的收敛速度明显更快。2.3 数据集划分与数据加载器实现数据划分的原则是训练集、验证集、测试集三层不能混淆。训练集用于学习参数验证集用于调参和选模型测试集用于最终评估。我按7:2:1的比例划分。这里有一个非常重要但容易被忽略的细节划分的时候一定要保证类别比例在每个集中保持一致也就是做分层采样。如果单纯随机划分很有可能某一类水果在测试集中数量特别少导致评估结果不真实。分层划分的代码from sklearn.model_selection import train_test_split all_images [] all_labels [] # 遍历所有类别生成图片路径和标签 for class_idx, class_name in enumerate(sorted(os.listdir(root_dir))): class_path os.path.join(root_dir, class_name) if not os.path.isdir(class_path): continue for img_name in os.listdir(class_path): if img_name.lower().endswith((.jpg, .jpeg, .png)): all_images.append(os.path.join(class_path, img_name)) all_labels.append(class_idx) # 先分出测试集 train_images, test_images, train_labels, test_labels train_test_split( all_images, all_labels, test_size0.1, stratifyall_labels, random_state42 ) # 再从剩余数据中分出验证集占原始数据的20% train_images, val_images, train_labels, val_labels train_test_split( train_images, train_labels, test_size0.222, stratifytrain_labels, random_state42 )使用PyTorch的DataLoader加载数据时我设置batch_size32shuffleTrue仅训练集num_workers4Windows系统上建议设为0或2设太高容易报错。shuffle的作用是让模型在每一轮迭代时看到的样本顺序都不同避免模型记住数据的排列顺序。3. CNN模型结构解析与选型权衡3.1 从零搭建简单CNN网络学习CNN最好的方式是自己动手搭一个简单的网络而不是直接套用现成的ResNet。这样你能直观理解卷积、池化、全连接这些核心概念到底是怎么运作的。这段代码我写得很克制就是纯粹的入门级结构import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes5): super(SimpleCNN, self).__init__() self.conv_layers nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(kernel_size2, stride2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(kernel_size2, stride2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(kernel_size2, stride2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(128 * 28 * 28, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) ) def forward(self, x): x self.conv_layers(x) x self.classifier(x) return x简单分析一下这个网络结构。输入是224x224x3的图片第一层卷积用32个3x3卷积核输出的特征图尺寸还是224x224因为padding1通道数从3变成32。ReLU激活函数的作用是引入非线性MaxPooling把特征图尺寸减半到112x112。经过三次卷积池化后特征图变成28x28x128。Flatten后进入全连接层先转成256维的向量再映射到5个类别。全连接层中间的Dropout(0.5)是一个很关键的技巧它会在训练时随机丢弃一半的神经元输出强迫模型不依赖于某个特定的神经元从而减少过拟合。这个网络的结构很经典运行速度也快在小数据集上训练的准确率大约能到70%-80%。作为理解CNN的起点个人强烈推荐先把这个网络完整跑通再去玩复杂网络。3.2 为什么最终我选择了ResNet简单CNN跑通以后我明显看到了它的瓶颈在测试集上的准确率一直在80%左右徘徊很难再往上升。主要原因在于它比较浅特征提取能力有限对于有些外观相近的水果比如青苹果和猕猴桃在颜色纹理上有一点点相似区分能力不足。后来我换成了ResNet-18。ResNet的核心创新是残差连接也就是说它允许信息跨层直接传递。在传统的CNN里每一层学习的是从输入到输出的完整映射函数网络加深以后很容易出现梯度消失——换句话说梯度传着传着就变成0了浅层参数很难得到有效更新。ResNet的做法是在网络中加入了一个捷径shortcut让每一层只需要学习输入和输出之间的残差难度降低了梯度也更容易流动。我为什么选ResNet-18而不是更深的ResNet-50或者ResNet-101原因很简单数据集规模不大总共才千来张图片ResNet-18有约1100万参数对这个数据量来说已经足够强了。强行上ResNet-50约2550万参数反而更容易过拟合而且训练速度也慢了很多。加载预训练权重的方式是import torchvision.models as models model models.resnet18(pretrainedTrue) num_features model.fc.in_features model.fc nn.Linear(num_features, 5)这里的关键操作是把ResNet自带的最后一层全连接层替换成输出维度为5的新全连接层因为原来的ResNet是在ImageNet的1000类上训练的不适用于我们5类水果分类的任务。num_features是最后一层之前的特征维度对ResNet-18来说是512。使用预训练权重本质上就是迁移学习预训练模型已经在ImageNet上学到了很多通用的视觉特征边缘、纹理、形状等我们只需要在水果数据上微调它。这个策略在数据量有限的时候非常有效可以让模型在小数据集上也能达到很高的准确率。我最终训练出来的模型测试准确率在96%左右远高于自己搭建的简单CNN。3.3 损失函数与优化器的选择逻辑损失函数我用的交叉熵损失CrossEntropyLoss这是多分类问题的标准选择。可以简单把它理解为如果一个样本的真实类别是苹果模型输出的每类概率分布是[0.6, 0.2, 0.1, 0.05, 0.05]那么这个预测和真实标签之间的交叉熵损失就会比较小如果模型输出的概率分布是[0.2, 0.3, 0.3, 0.1, 0.1]类别比较分散损失就会大。训练的目标就是让这个损失值不断下降。优化器我选的是Adam学习率初始设成1e-4。这里有个细节值得展开讲讲就是用预训练模型做迁移学习时不同层的学习率应该不一样。因为预训练层已经收敛过了应该用较小的学习率去微调避免破坏已经学到的特征而新加的全连接层是从零开始训练的需要较大的学习率才能快速收敛。实际代码里我通过给不同参数组设置不同的学习率来实现optimizer torch.optim.Adam([ {params: model.conv1.parameters(), lr: 1e-5}, {params: model.layer1.parameters(), lr: 1e-5}, {params: model.layer2.parameters(), lr: 1e-5}, {params: model.layer3.parameters(), lr: 1e-5}, {params: model.layer4.parameters(), lr: 1e-5}, {params: model.fc.parameters(), lr: 1e-4} ])不过我训练到后面发现全数据集统一用1e-4也能得到不错的效果因为我们的数据量不大而且和ImageNet的自然图像分布很像微调不需要做太复杂的差异化学习率设置。新手用简单的统一学习率就够了。学习率调度我用的StepLR每隔10个epoch把学习率乘以0.1这样训练后期模型能更精细地收敛。4. 训练全流程从损失下降到模型收敛4.1 训练循环的具体实现与监控训练一个CNN模型的核心就是一个循环把数据一批一批地喂给模型计算损失反向传播更新参数。但真正动手写的时候有一些细节处理会直接决定你训练过程是否顺利。我最开始用的训练循环代码如下def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for images, labels in dataloader: images images.to(device) labels labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc correct / total return epoch_loss, epoch_acc这里有几个细节值得新手注意optimizer.zero_grad()每次迭代开始前必须调用否则梯度会在反向传播时累加导致参数更新方向错乱。我在网上看到过一些教学代码漏掉这一步然后跑来问我为什么模型不收敛十有八九就是这个原因。model.train()这个模式切换也很重要。因为在训练时Dropout和BatchNorm等层的行为和推理时不同——Dropout在训练时会随机丢弃神经元BatchNorm会使用当前batch的统计信息。如果用训练模式做推理预测结果会不稳定。所以验证时一定要调用model.eval()并用torch.no_grad()关闭梯度计算因为验证阶段不需要反向传播。训练过程中实时监控的指标主要是两个训练集loss和验证集loss。每跑完一个epoch我都会打印出来看看并且保存一版当前状态。判断训练是否正常的标准是训练集loss持续下降验证集loss也同步下降说明模型在正常学习如果训练集loss在下降但验证集loss在上升说明过拟合了需要加数据增强或Dropout如果训练集loss和验证集loss都停滞不动那可能是学习率太大或太小模型参数根本没在有效更新。4.2 过拟合观察与早停机制训练大概到第20个epoch时我观察到训练集准确率已经接近100%但验证集准确率在95%左右徘徊不动。这是典型的过拟合信号——模型开始背答案了把训练集里的某些特殊背景、光影条件记住了而不是真正学会识别水果本身。针对过拟合我采用了几个措施第一个是加强数据增强。把原先的随机旋转幅度从15度扩到30度再加上RandomResizedCrop让模型每次看到的都是同一张图的不同区域和不同缩放比例这样它就很难记住图片的细节信息。第二个是引入早停Early Stopping。这是我觉得最实用的机制它的逻辑是每轮训练完在验证集上评估如果验证损失连续多个epoch我设置的是10个epoch没有下降就停止训练并且回滚到验证损失最低的那个模型权重。这样能避免在过拟合区域里白费时间和算力。class EarlyStopping: def __init__(self, patience10, min_delta0.001): self.patience patience self.min_delta min_delta self.counter 0 self.best_loss None self.should_stop False self.best_model_state None def __call__(self, val_loss, model): if self.best_loss is None: self.best_loss val_loss self.best_model_state model.state_dict().copy() elif val_loss self.best_loss - self.min_delta: self.best_loss val_loss self.best_model_state model.state_dict().copy() self.counter 0 else: self.counter 1 if self.counter self.patience: self.should_stop True return self.should_stop实际训练中模型大约在第18个epoch时达到最佳验证损失之后虽然训练集损失还在下降但验证损失不再改善甚至上升。早停机制把模型的最终权重定格在了第18轮的版本准确率也是整个训练过程中最高的。第三个措施是调整Dropout的比例。我尝试过把全连接层前的Dropout从0.5提高到0.6但对这个数据集来说效果差别不大所以最后就保留了0.5。4.3 训练时长与GPU资源消耗这一步我单独拿出来讲因为很多初学者对训练需要多少时间、多少显存没有概念。训练时的一个意外情况是这样的我的RTX 3060 12GB显存实际上非常充裕因为数据集小、图片尺寸也小整个模型加数据所占用的显存大约只有2GB左右。很多人以为深度学习非要顶级显卡不可其实对于这种小规模的图像分类项目普通显卡完全够用。训练过程解析如下数据加载阶段CPU负责读图片、做数据增强、把图片转成tensor并放到GPU上这个阶段大约占总时长的30%左右前向传播阶段GPU计算每张图片经过卷积网络后的输出这个阶段大约占30%后向传播阶段计算梯度并更新参数这个阶段占40%左右完整训练一轮遍历所有训练集图片大约需要25秒30个epoch总共不到15分钟。这个速度比我预期的快很多因为数据量小、网络结构也不是特别深。如果在训练时感觉速度太慢可以按这个顺序排查先看是不是没有用GPU用CPU训练会慢几十倍再看是不是num_workers设得太低导致CPU在读图片时GPU在空转最后看是不是batch_size设得太大导致显存不够反而触发swap拖慢速度。4.4 训练过程中的Loss曲线分析训练结束后我把训练集loss和验证集loss随epoch变化的曲线画了出来。这个可视化步骤强烈推荐大家做它能直观地反映训练健康度。在我的训练中前几个epoch训练集loss从2.2左右快速下降到0.8左右这说明模型在快速学习区分不同水果。到第10个epoch时训练集loss降到了0.3左右验证集loss在0.25左右。第18个epoch时验证集loss到达最低点0.18之后开始轻微回升。训练集loss在差不多第15个epoch时就已经降到了0.05以下继续训练其实意义不大了。这种训练集和验证集loss前期同步下降、后期训练集继续下降但验证集停滞的曲线形态在深度学习里非常经典。它说明模型的能力已经达到了数据量支撑的上限再多的训练轮次也不会带来提升应该考虑的是增加数据量、扩大数据增强力度或者使用更高效的网络结构。5. 模型评估与混淆矩阵分析5.1 混淆矩阵看分类错误训练完之后我用测试集评估模型表现。整体准确率是96.3%对于5类水果分类来说这个数字已经不错了但我觉得还应该进一步看错误到底出在哪里。用混淆矩阵来分析是最直观的方式。我先是把每个类别的召回率和精确率都算了一遍from sklearn.metrics import classification_report, confusion_matrix y_true [] y_pred [] model.eval() with torch.no_grad(): for images, labels in test_loader: images images.to(device) labels labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) y_true.extend(labels.cpu().numpy()) y_pred.extend(predicted.cpu().numpy()) print(classification_report(y_true, y_pred, target_names[apple, banana, grape, orange, kiwi])) print(confusion_matrix(y_true, y_pred))输出的报告显示大部分类别的精确率和召回率都在95%以上但葡萄和猕猴桃之间有一些混淆。具体原因是有些青色葡萄的图片在形状和颜色上非常接近猕猴桃切面的纹理模型判断错了我觉得也可以理解。有一个让我印象深刻的错误是有一张图片拍的是香蕉的尾部特写颜色偏绿模型把它预测成了青苹果。这说明模型的判断依据主要还是颜色和形状特征对于这种最容易被颜色干扰的图片想要进一步区分率可能需要加入更多纹理特征或者提升数据集里类似样本的数量。5.2 测试不同图片的预测效果我用一张从未参与训练的真实水果照片来测试模型的泛化能力。这张照片是我自己用手机拍的光线条件、背景、拍摄角度都与数据集里的图片差异很大。预测过程的代码很简单from PIL import Image import torchvision.transforms as transforms def predict(image_path, model, device, class_names, mean, std): transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(meanmean, stdstd) ]) image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): output model(input_tensor) probabilities torch.softmax(output, dim1) confidence, predicted_class torch.max(probabilities, 1) return class_names[predicted_class.item()], confidence.item()这里有个细节挺重要预测时用的预处理流程必须和训练时完全一致不然图片分布的差异会直接影响预测结果。比如训练时用了RandomResizedCrop但预测时就应该用CenterCrop这是固定的操作。我拍的照片是一根切好的香蕉放在白色盘子里背景是木桌。模型以98.7%的置信度正确识别为香蕉。另外我又找了一张网上下载的葡萄串图片背景复杂得多模型还是以93.2%的置信度成功识别。这次泛化实验结果给我挺大信心的说明模型学到的是水果本身的视觉特征而非训练集中背景的作弊特征。6. 部署与模型导出6.1 保存与加载模型的最佳实践训练完成后我把模型保存了下来。PyTorch中保存模型有两种方式新手很容易混淆第一种是保存整个模型第二种是只保存模型的状态字典。推荐使用第二种因为它更灵活而且不会因为代码版本迭代导致加载失败。# 保存 torch.save({ model_state_dict: model.state_dict(), class_names: class_names, mean: mean.tolist(), std: std.tolist(), }, fruit_classifier.pth) # 加载 checkpoint torch.load(fruit_classifier.pth) model.load_state_dict(checkpoint[model_state_dict])把类名、均值、标准差都存进同一个文件是一个很实用的习惯。否则你部署的时候还要额外维护一个配置文件并且容易忘记训练时用的预处理参数是多少。6.2 CPU环境下的推理部署实际部署不一定有GPU可用我特意测了一下只在CPU上推理的速度。用PyTorch自带的torch.jit做一下脚本化转换然后仅用CPU推理识别一张图片大约需要80-120毫秒速度完全够用。如果还想更快可以量化转换成半精度或用ONNX导出后者可以在不同的推理框架上运行。不过这个项目规模不大用PyTorch原生推理就够了暂时不需要引入额外的部署复杂度。6.3 模型裁剪思路有没有更轻量的方案如果在资源更受限的嵌入式设备上部署ResNet-18可能还是太大。我做一个简单估算模型权重文件大小约45MB内存占用约90MB因为推理时还会产生中间特征图。如果目标设备只有几十MB的内存可以考虑MobileNetV3或ShuffleNetV2这类轻量网络它们的参数量只有ResNet-18的几分之一精度损失在1-2个百分点以内。使用轻量网络的方法其实很简单只需要把models.resnet18(pretrainedTrue)换成models.mobilenet_v3_large(pretrainedTrue)然后把最后的全连接层输出改成5类其他训练流程不需要改动。这也是建议深度学习入门者把网络结构封装成单独函数的原因换网络结构就是一行代码的事。7. 踩坑记录与常见问题排查7.1 图片加载报错PIL无法识别文件第一次训练时我在数据加载阶段遇到一个报错PIL.UnidentifiedImageError: cannot identify image file。排查后发现数据集里有几张图片后缀是.jpg但实际是损坏的文件Image.open()虽然不报错但load()的时候撑不住了。我加了一个try-except把这类图片过滤掉同时删除了全部损坏文件。如果你在训练过程中也遇到这个报错建议先写一个脚本遍历所有图片验证完整性而不是逐个排查太费时间。7.2 验证集准确率为0的问题这个问题发生在我第一次运行验证代码时因为忘了调用model.eval()而且在验证循环里没有加torch.no_grad()导致Dropout层在验证时仍然随机丢弃神经元输出结果一会儿一变。有时候准确率看起来还行有时候所有预测都集中在某一个类别上看起来就像准确率为0。解决方法很简单在验证前加上model.eval()以及用with torch.no_grad():包裹前向传播代码。这个坑几乎每个新手都会踩建议直接形成肌肉记忆。7.3 GPU显存不足或训练速度极慢对于这个项目RTX 3060训练毫无压力但我试过把batch_size调到256结果直接OOMOut of Memory了。根本原因是每个batch的图片同时驻留在显存里计算batch越大中间特征图占用的显存就越多。如果显存不够优先降低batch_size然后适当降低图片分辨率再不行就用梯度累积的方式模拟大batch训练。但就这个数据集而言batch_size32已经足够了再大训练速度也不会显著提升。7.4 训练不收敛损失函数一直不下降我也遇到过训练loss一直不下降的情况原因是学习率设置得太大了我一开始用了0.01。学习率太大时参数更新步长太大loss在某个范围内震荡甚至发散。解决办法是换小一点的学习率比如从0.01降到0.001或0.0001。我的经验是用Adam优化器时初始学习率设为1e-3通常是一个稳妥的选择如果用预训练模型微调用1e-4更安全。8. 项目总结与扩展方向这个项目花费了大约两天的时间中间踩了不少坑。回头看最大的收获不是跑通了一个模型而是彻底理解了CNN图像分类的完整流程从数据清洗开始到预处理、模型设计、训练调优、评估分析、最终部署每个环节都有大量可以优化的细节。如果大家想做进一步扩展可以按照这个思路往下走增加类别数量从5类扩展到几十类观察准确率下降的幅度然后针对性做数据增强换更强的网络结构尝试ResNet-50或EfficientNet对比效果差异加入注意力机制在CNN里加入SE模块或CBAM模块让模型更关注包含关键特征的重要通道用真实场景图像做测试从网络上爬取大量真实场景的水果图片测试模型的泛化能力根据我这次实际操作的体会最重要的一条经验是深度学习模型的效果很大程度取决于你如何处理数据而不是你用了多深多复杂的网络。把数据的每个细节都处理好一个简单的ResNet-18就能在水果分类任务上达到95%以上的准确率反之即使你上了ResNet-200数据脏乱差也一样训练不出理想效果。希望这篇完整的过程记录能给想入门CNN图像分类的朋友提供一条可以全程顺利跑通的路。本文还有配套的精品资源点击获取
返回列表