简介:基于深度残差网络(ResNet)的水果分类识别系统完整代码包,面向具备一定Python基础、希望快速落地图像分类项目的开发者与学生,尤其适合需要完成课程设计、毕业设计或工程演示的入门者。项目以水果分类为例,覆盖数据预处理、TFRecord生成、TensorFlow模型构建与训练、预测评估全流程,核心代码可直接复用,更换图像与标注文件即可适配其他分类场景。资源共8780个文件,以8767张jpg训练/测试图片为主,附3个ipynb教学笔记、模型权重checkpoint及Python脚本,压缩包约588MB,便于离线学习与调试,目录组织也便于按数据处理、训练、评估三个环节分别查阅。已有6500人浏览学习,适合作为入门深度残差网络的实战参考。借助该代码包,读者可获取完整的ResNet分类实现思路与可运行代码,无需从头推导复杂原理即可快速跑通流程,省去大量环境配置和排错时间,尤其适合需要快速产出可演示系统或进行二次开发的场景。 相信不少朋友都遇到过这种尴尬:翻遍各种教程,装好了环境也跑通了代码,模型在验证集上表现得漂漂亮亮,可一到自己拍的照片就现出原形——把青苹果认成青梨,把熟透的香蕉认成芒果。如果你正准备做目标识别或者说水果分类这类入门级视觉项目,本文应该能帮你避开我当初踩过的大部分坑。我会用一套基于深度残差网络(ResNet)的完整水果分类识别系统,把从数据集准备、模型搭建到训练调参、评估部署的完整链路拆开揉碎讲清楚。
这套系统基于 PyTorch 实现,是一个标准的多分类图像识别任务:输入一张水果图片,模型输出它属于哪种水果。听起来简单,但要把准确率做到可用级别、让模型真正具备泛化能力,里面涉及的技术细节远比想象中多。适合刚入门计算机视觉、准备做课程设计或者想系统梳理一遍图像分类流程的开发者参考,也适合那些已经跑通过分类模型、但总感觉自己对整个流程缺乏整体把控的朋友查漏补缺。
1. 数据准备与预处理:分类系统的地基工程
1.1 数据集选型:为什么我选了 Fruits-360
水果分类项目最怕的就是数据随便凑。我最早尝试过自己拍照建数据集,结果苹果在不同光照下拍出来的颜色差异比不同品种之间的差异还要大,模型训练出来直接崩溃。后来换了公开数据集 Fruits-360,这是一个专门为水果识别任务制作的数据集,目前已经包含上百类水果,每类几百到上千张不等,全部为白底图,单个水果居中摆放。
这个数据集最大的优点不在数量,而在类别的清晰划分。它把不同成熟度、不同品种的苹果、梨、香蕉分成了独立类别,比如 GreenApple、RedApple 1、RedDelicious 等,这对训练一个严谨的分类器来说非常重要——如果类别内部差异过大,模型会无所适从。
1.2 数据划分的关键坑:同源图片泄露
数据准备阶段最容易被忽视但杀伤力极大的问题,就是同源图片的数据泄露。Fruits-360 每个类别下的图片实际上是同一个水果在旋转台上旋转不同角度拍摄的连续帧,如果不加处理直接随机切分训练集和验证集,同一水果的不同帧会同时出现在两边,验证集的准确率会被严重虚高。
我当时的处理方式是:按照图片文件名中的 ID 进行分组,确保同一个水果的所有帧只进入训练集或只进入验证集,绝不跨集。分组完再按 8:1:1 切分为训练集、验证集和测试集。这一步做完之后,我的验证集准确率从虚高的 99% 回落到真实的 96% 左右,两者差距直观反映了数据泄露问题的严重性。
- 逐张随机切分 → 同源帧泄露 → 验证集虚高,上线后崩溃
- 按水果 ID 分组切分 → 真实泛化能力 → 模型更可靠
1.3 数据增强策略:不要盲目堆砌
很多教程喜欢把 ColorJitter、RandomRotation、RandomResizedCrop 全部堆上,看得人热血沸腾,实际跑起来却发现训练怎么都不收敛。这里的关键认知是:数据增强的本质是模拟真实世界的变异,而不是单纯增加数据量。
对水果识别来说,我认为最有效且稳妥的是这四板斧:
train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])RandomResizedCrop 模拟不同拍摄距离;RandomHorizontalFlip 模拟左右摆放角度变化;ColorJitter 的幅度控制在 0.2 是因为水果颜色本身就是分类的关键特征,调过头会破坏语义;Normalize 用的是 ImageNet 的均值和标准差,因为后续要用预训练权重做迁移学习,输入分布必须对齐。
验证集和测试集只做 Resize 到 256、CenterCrop 到 224,不做任何随机增强,保证评估的确定性。
2. ResNet 核心机制拆解:为什么加深网络反而会退化
2.1 退化问题与残差学习的动机
在 ResNet 出现之前,大家普遍认为网络越深表达能力越强,但实验很快打了脸。一个 56 层的普通卷积网络在 CIFAR-10 上的训练误差居然高于 20 层的版本,这不是过拟合,因为训练误差本身就更高——说明深层网络在优化层面就出了问题。
原因出在恒等映射难以学习。一个 20 层的网络理论上可以模拟出 56 层网络中前 20 层学到的东西、后 36 层什么都不做(恒等映射),但如果让普通卷积层直接拟合恒等映射,权重矩阵根本不容易收敛到单位矩阵。你我在实践中更直观的感受是:网络越深,梯度在反向传播中连乘后趋于消失,浅层参数根本得不到有效更新。
2.2 残差模块的数学直觉
ResNet 的解决方案是显式地在网络结构中加入一条"短路"分支,让梯度有一条畅通的通道回传。残差模块的计算可以写成:
$$y = \mathcal{F}(x, {W_i}) + x$$
其中 $\mathcal{F}$ 代表卷积层要学习的残差映射,$x$ 是输入,$y$ 是输出。当网络觉得当前层没有新特征需要提取时,只需要让 $\mathcal{F}$ 的输出逼近 0,$y$ 就等于 $x$,恒等映射变得非常容易实现。
这个过程可以做一个直观类比:普通网络就像一个人闭着眼睛记忆每天的路线,任何一天走错了都会越偏越远;ResNet 则像在沿途钉了路标,允许偏差被随时拉回主路。即使中间某层没学好,梯度也能通过旁路直接回流到浅层,解决了深层网络训练的根本障碍。
2.3 残差结构代码实现
这是 ResNet 中最基础也最重要的 BasicBlock 实现,对应 ResNet18/34 的结构,两个 3x3 卷积组成一个残差块:
import torch.nn as nn class BasicBlock(nn.Module): def __init__(self, in_channels, out_channels, stride=1): super().__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_channels) self.shortcut = nn.Sequential() if stride != 1 or in_channels != out_channels: self.shortcut = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity = self.shortcut(x) out = torch.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) out += identity return torch.relu(out)注意两点:一是当 stride 不为 1 或通道数改变时,shortcut 分支需要加一个 1x1 卷积来做维度匹配;二是 BatchNorm2d 在 Conv 之后、激活之前,这是 ResNet 的标准排列。
2.4 结构选型:为什么 ResNet34 更适合水果分类
ResNet 家族里,我最终选择的不是 18 也不是 50,而是 34。18 层太浅,对水果纹理、颜色、形状的联合特征提取能力有限;50 层引入了 Bottleneck 结构,参数数量呈指数增长,但对水果这种并不是极度复杂的分类任务来说,收益边际递减,训练和推理成本却显著上升。
ResNet34 的每层通道数设计为 64、128、256、512,通过四个 stage 逐级扩大感受野并缩减特征图分辨率,最终经过全局平均池化输出 512 维特征向量,再接一个全连接层映射到类别数。这里有个容易被忽略的设计细节:全局平均池化替代了传统 Flatten + 全连接,这大大减少了参数数量,天然具备正则化效果,这也是 ResNet 结构设计远比 VGG 精简的核心原因。
3. 完整实现:从数据加载到训练流程
3.1 自定义数据集类
虽然 Fruits-360 的目录结构是标准的按类别分文件夹,我仍然建议实现一个自定义 Dataset,把图片路径和标签的映射显式管理起来,方便后面的分组切分和后续换数据集复用。
import os from PIL import Image from torch.utils.data import Dataset class FruitDataset(Dataset): def __init__(self, root_dir, class_to_idx, transform=None): self.samples = [] self.transform = transform for class_name, idx in class_to_idx.items(): class_dir = os.path.join(root_dir, class_name) for img_name in os.listdir(class_dir): self.samples.append((os.path.join(class_dir, img_name), idx)) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] image = Image.open(img_path).convert('RGB') if self.transform: image = self.transform(image) return image, labelclass_to_idx 可以从训练集目录中按字母序生成,保证训练、验证、测试三者的类别编码完全一致。我记得一开始图省事直接用了 PyTorch 自带的 ImageFolder,结果在按 ID 分组切分数据时被目录结构绑死,折腾了半天才改写成现在这个版本。
3.2 迁移学习:加载预训练权重并解冻策略
直接随机初始化 ResNet34 从头训练,在水果这种中等规模数据集上效果很一般,收敛也慢。更靠谱的做法是加载 ImageNet 上预训练好的权重,利用它在海量自然图像上学到的通用特征——边缘、纹理、颜色分布——作为起点。
import torchvision.models as models model = models.resnet34(weights=models.ResNet34_Weights.IMAGENET1K_V1) num_features = model.fc.in_features model.fc = nn.Linear(num_features, num_classes) # 锁定前三个 stage 的参数,只微调最后一个 stage 和全连接层 for name, param in model.named_parameters(): if 'layer4' not in name and 'fc' not in name: param.requires_grad = False这里的关键是解冻策略。我的经验是:第一阶段冻结大部分参数、只训练最后一个 stage 和全连接层,让新分类头先稳定下来;训练一段时间后再解冻所有参数,用很小的学习率整体微调。这两个阶段的学习率通常差一个数量级。
3.3 训练脚本的工程化细节
训练循环本身不复杂,但有几个工程细节能显著提升体验:
import torch def train_one_epoch(model, train_loader, criterion, optimizer, device): model.train() running_loss, correct, total = 0.0, 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() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() running_loss += loss.item() * images.size(0) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() return running_loss / total, correct / total梯度裁剪clip_grad_norm_是我在训练初期 loss 偶尔爆掉之后加的保险,尤其是解冻所有层进行全模型微调时,预训练参数的大梯度很容易把已经学好的特征摧毁。优化器我用的是 AdamW,初始学习率 1e-3,权重衰减 5e-4,配合余弦退火学习率调度器,整体训练曲线比固定学习率平滑很多。Batch size 设为 32,在单张 8G 显存的显卡上训练 ResNet34 毫无压力。
4. 训练过程中的实际问题与排查链路
4.1 诡异案例一:验证集准确率卡在 93% 不涨了
这是我第一次训练跑到第 40 个 epoch 时遇到的瓶颈。训练准确率已经接近 99%,验证集却一直卡在 93% 上下,典型的过拟合信号。排查链路是这样的:
第一步检查数据增强的强度,把 ColorJitter 的幅度从 0.2 降到 0.1,同时加上 RandomRotation(10),模拟水果摆放角度的轻微偏差。第二步检查模型结构,确认最后的全连接层后是否忘了加 Dropout——在 fc 层前加了一个 p=0.3 的 Dropout。第三步是降低学习率,把 1e-3 降到 3e-4,让模型在损失曲面更精细的区域里继续探索。
三步同时做之后,验证集准确率在 10 个 epoch 内突破到 96.5%。这让我意识到过拟合的成因往往是多方面的,单点修复通常不够,需要组合拳。
4.2 诡异案例二:模型把"青苹果"系统性误分类为"青梨"
这是一次典型的类间相似度过高导致的语义混淆。青苹果和青梨在形状、颜色、纹理上的差异确实很小,连人也容易看错。我通过混淆矩阵发现,两类之间的误分类贡献了整个错误率的三分之一。
一般的解决思路是:收集更多针对这两类的数据,或者引入额外的判别特征。但我当时没有更多数据可用,于是做了一个很有效的调整:把标签从单一分类改成加入一个"辅助分类头"——主分类头输出具体水果类别,辅助分类头输出水果的高层属性(如柑果类、梨果类、浆果类等)。多任务学习迫使共享特征在保留类别细节的同时,提取到更高层的语义共性,最终这两类的混淆率明显下降。
如果你不想引入多任务结构,一个更简单的办法是在损失函数上做文章,加大困难样本的权重,使用 Focal Loss 替代普通的交叉熵损失,让模型把注意力放在这类难分样本上。
4.3 训练状态监测:loss 曲线和梯度范数
把训练过程变成可视化曲线之前,调参基本靠玄学。我后来在训练脚本里每 50 步记录一次 loss 和梯度范数,并同步记录学习率变化,这样才能准确定位问题:
- loss 震荡剧烈且梯度范数超过 10:大概率是学习率过高,建议下降一个数量级
- loss 平滑下降但验证集不涨:模型欠拟合,需要加大容量或解冻更多层
- 训练 loss 快速降到接近 0 而验证集很差:过拟合信号,优先调数据增强和 Dropout
- 二次学习率重启后 loss 反弹:余弦退火的周期和总 epoch 数不匹配
WandB 和 TensorBoard 都行,我倾向于用 TensorBoard,零成本接入。关键是养成看曲线的习惯,而不是只盯最后那个准确率数字。
5. 评估与部署:从准确率到真正可用
5.1 混淆矩阵和单类别指标比总体准确率更重要
水果分类这种类别数较多的任务,总体准确率会掩盖个别类别的失效。我最终排查青苹果问题时,靠的就是混淆矩阵,这一点在第 4.2 节已经提到。
测试完成后,我还会额外统计每个类别的精确率、召回率和 F1-score。比如模型对芒果的召回率只有 85%,意味着 15% 的芒果被漏掉了,可能是某些品种的芒果颜色偏绿,训练集中占比太少。针对这一类样本做简单的过采样,比整体加数据更有效。
5.2 部署前必须检查的数据一致性
模型训练完成,封装成推理接口只花了半天,真正花时间的是排查一个诡异现象:训练时准确率 96%,部署到本地跑一张测试图片,结果怎么都是错的。
最后定位到的原因极其基础:训练时输入经过了 Normalize,而推理脚本里忘了做同样的预处理。另外还有一个隐蔽的坑是训练时用的是 Resize 到 256 再 CenterCrop 224,但推理时用了直接 Resize 到 224,导致图片的比例和感受野分布都变了,模型自然就罢工了。
建议在部署代码里把数据预处理单独抽成一个函数,跟训练代码共用同一份实现,从根源上避免这种低级但致命的不一致。
5.3 导出模型并封装推理函数
训练结束,把权重保存成完整模型文件,同时保留一份只含 state_dict 的版本,方便后续迁移。推理时用 torch.jit 或 ONNX 导出做加速,这一步对移动端部署尤其重要:
import torch import torchvision.transforms as transforms from PIL import Image model.load_state_dict(torch.load('fruit_resnet34.pth', map_location='cpu')) model.eval() example = torch.randn(1, 3, 224, 224) traced_model = torch.jit.trace(model, example) traced_model.save('fruit_resnet34_jit.pt')使用 TorchScript 导出的模型不再依赖原始的 Python 类定义,部署到服务端甚至嵌入式设备时省心很多。推理时如果对置信度低于 0.7 的结果统一返回"未知水果",能过滤掉大量模型不确定的输入,整体体验会好很多。
6. 进一步优化方向
从这套水果分类系统出发,有很多可以继续深入的方向。最简单的改进是增加类别数量,目前 Fruits-360 已支持百级类目,把 ResNet34 换成 ResNet50 并加入更复杂的增强策略,准确率还能继续往上走。
如果想让模型具备更强的细粒度识别能力,可以引入注意力机制模块,比如在最后一个 stage 后接入 SE Block 或 CBAM,让模型自动聚焦于水果的局部判别区域。对移动端或嵌入式场景,可以考虑轻量化网络如 MobileNetV3 或 EfficientNet-Lite,配合知识蒸馏把 ResNet34 学到的知识迁移到轻量模型上,在几乎不掉点的条件下将推理速度提升数倍。数据层面,如果后续有采集条件,补充不同光照、不同背景、部分遮挡的自然场景图片往往比单纯增加白底图更有价值,这也是让模型从实验室走向真实环境的最关键一步。
本文还有配套的精品资源,点击获取