ARTICLE DETAIL

资讯详情

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

小样本果蔬图像分类实战:基于PyTorch的迁移学习与部署指南

小样本果蔬图像分类实战:基于PyTorch的迁移学习与部署指南 简介面向图像分类算法研究与课程实验这份资源提供了36类常见水果和蔬菜的已标注图像数据集包含约3400张图片并按训练集、验证集分目录存放覆盖香蕉、苹果、番茄、茄子等日常类别。数据经过预处理可直接作为分类网络输入省去自行采集、清洗与标注的时间适合刚接触计算机视觉的读者快速上手。资源包共2000个文件主体为1998张图片另附1个Python可视化脚本和1个JSON类别映射文件整体约94.47MB解压后目录结构清晰按训练集、验证集组织。脚本支持随机查看各目录样本图片JSON文件则保存36类标签与名称对应关系可辅助核对数据划分是否正确。目前已有212人学习下载这套数据可直接用于模型训练、精度对比或课程设计是一份省心易用的基准数据集。1. 3400 张图像做 36 类分类先从数据集容量读起一份“已标注、约 3400 张”的水果蔬菜图像数据集平均到 36 个类别每类只有 94 张左右。这个量级很尴尬比 MNIST、CIFAR 这类教学数据小一个数量级又比工业里常见的“每类几百上千张”少很多。放在 2024 年再往后看的背景下它既不够喂给大模型微调又足以撑起一个可用的图像分类基线任务。真正值得做的不是“训练一个模型”而是把数据组织、增强策略、迁移学习、验证曲线这一整套小样本图像分类的工程路径走通。这个数据集能解决两类问题一是给图像分类入门提供一个不烧显卡的训练样本二是给“类别多、单类样本少”的真实业务打底——农产品质检、商超称重、拍摄设备无关的果蔬识别场景本质都是如此。适合的人群包括刚接触 PyTorch 的开发者、需要做 PoC 验证的算法工程师以及想在自采数据上复制同样流程的团队。整篇文章我会按自己拿到这份数据后会做的事来展开先确认标注格式再把训练管线搭起来然后处理增强和调参最后落到验证与导出。2. 数据集结构与标注格式先数清楚 36 类是什么拿到“约 3400 张已标注”的数据第一件事不是写模型而是确认标注是以什么形式存在的。常见情况有三种按类别建文件夹、单张图像配一个标签文件、或者是 COCO/LabelMe 风格的 JSON 标注。36 类果蔬的公开数据集大多走前两种因为图像里通常只有一个主体目标语义分割级别的标注对图像分类任务反而是多余的。2.1 类别分布与每类样本量预估3400 除以 36每类约 94 张但这只是平均值。实际数据集的类别分布经常不均衡比如香蕉、苹果这类常见水果可能有一百五六十张而某些叶菜可能只有六七十张。这种不均衡直接决定了后面要不要做类别加权采样所以先用一个命令把各类别样本数统计出来。以按文件夹分类的目录结构为例Linux 或者 macOS 下直接数文件# 统计 train 目录下每个子文件夹的图片数量 # find 递归查找-type f 只统计文件cut 截取类别目录名sortuniq 聚合计数 find train -type f \( -name *.jpg -o -name *.png -o -name *.jpeg \) \ | cut -d/ -f2 \ | sort \ | uniq -c \ | sort -rncut -d/ -f2这里假设目录结构是train/类别名/文件名.jpg按斜杠切分后取第二段就是类别名。如果数据目录层级更深比如train/一级类/二级类/文件名需要把-f2改成对应字段。统计出来的数值低于 50 的类别要标记出来这类样本在后续划分验证集时每类只剩 10 张左右很容易让验证集上的准确率出现几个百分点的随机波动。建议把这个统计结果整理成一张类别映射表格式固定下来后面所有脚本都复用这一份类别索引类别目录名原始样本数划分后训练数划分后验证数0apple156124321banana14211329...............35zucchini715615类别索引的顺序来源于sorted(os.listdir(train_dir))的排序结果不是人为拍脑袋定的。PyTorch 的ImageFolder会自动按文件夹名的字典序生成索引如果之后要导出模型给其他工程用这个顺序必须冻结成一份 JSON 存下来否则推理时类别对不上。2.2 目录式标注与 JSON 标注的取舍如果拿到的是每个类别一个文件夹的结构可以直接用torchvision.datasets.ImageFolder加载不需要写额外的标注解析代码。这是最快路径。但需要注意ImageFolder 不校验图像能不能被解码损坏的图片会在训练中途触发PIL.UnidentifiedImageError所以第一轮预处理时最好把不可读文件先揪出来# 用 find 找出所有图片文件逐个交给 file 命令检查 MIME 类型 # grep -v 过滤掉正常的 image/jpeg 和 image/png剩下的就是异常文件 find . -type f \( -name *.jpg -o -name *.png \) -print0 \ | xargs -0 file \ | grep -v image/jpeg \ | grep -v image/png如果file命令输出里出现data、text或者其他非 image 开头的类型对应文件大概率已经损坏或者后缀名与真实格式不符直接从训练目录里挪走。这一步对“约 3400 张”这种规模的数据集花不了两分钟但能避免训练跑到一半才崩。另一种常见组织方式是 CSV 或 JSON 标注每行记录image_path,label。这种情况我会优先转成文件夹结构而不是写自定义 Dataset原因是 36 类果蔬的图像分类任务用不到复杂标注字段目录结构可以直接复用 torchvision 的现成接口后续做类别筛选、合并、重命名都更方便。转换脚本只需要读 CSV、遍历每一行、把文件复制到target/label/目录下不需要额外贴代码。3. 用 PyTorch 搭最小训练管线ImageFolder 预训练模型数据格式确认完毕下一步就是让模型跑起来。我在这类小规模图像分类任务里的固定做法是加载一个在 ImageNet 上预训练过的卷积网络把最后的全连接层换成 36 类输出先用冻结主干的方式跑几个 epoch 看数据有没有问题再决定要不要解冻全量微调。3.1 定义数据变换与 DataLoader3400 张图全部塞进显存不现实也没必要。采用标准做法训练尺寸设为 224x224验证尺寸也设为 224x224但训练集和验证集用不同的变换组合。训练时做随机裁剪和水平翻转验证时只做缩放和居中裁剪保证验证指标不被人为增强干扰。import torch from torchvision import datasets, transforms # 训练集变换随机裁剪加水平翻转ImageNet 均值和标准差做归一化 # RandomResizedCrop 的 scale 范围控制在 0.6~1.0避免裁剪后丢失太多主体 train_transform transforms.Compose([ transforms.RandomResizedCrop(size224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(p0.5), 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]) ]) # ImageFolder 会按目录名自动生成类别索引 train_dataset datasets.ImageFolder(data/train, transformtrain_transform) val_dataset datasets.ImageFolder(data/val, transformval_transform) train_loader torch.utils.data.DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue ) val_loader torch.utils.data.DataLoader( val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue )这里shuffleTrue只对训练集打开验证集保持顺序是为了和类别索引对齐便于后续画混淆矩阵。num_workers4在本地机器上够用如果数据读取成为瓶颈也就是 GPU 利用率持续低于 70%可以逐步加到 8。pin_memoryTrue配合 GPU 训练能减少一次内存拷贝虽然数据量小收益不大但养成习惯没有坏处。这里要特别留意RandomResizedCrop的scale(0.6, 1.0)。默认值是(0.08, 1.0)那是为 ImageNet 那种主体占比悬殊的大规模数据集设计的。果蔬图像通常主体突出、占比很高裁剪比例太激进会把果实切掉一半反而降低分类精度。这是我在这类小数据集上踩过最多次的增强参数坑。3.2 Baseline 模型选择与训练循环主干网络我一般先从 ResNet18 开始。原因是 3400 张图对一个 1100 万参数的模型来说已经接近容量上限ResNet50 这类更深的网络在小数据上更容易过拟合而且训练产出的收益很小。如果追求更高精度第一步不是换更大的主干而是先把手头数据的增强和调参做足。换成 EfficientNet-B0 或者 MobileNetV3 也是合理的它们参数量更小但训练时需要额外留意 BatchNorm 的统计量更新问题。import torch.nn as nn import torch.optim as optim from torchvision import models # 加载 ImageNet 预训练权重useTrue 表示下载预训练参数 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # ResNet18 的最后一层是形如 (512, 1000) 的 Linear 层 # in_features 取原全连接层的输入维度替换成 36 类输出 num_classes 36 in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-3) def train_one_epoch(model, loader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for images, labels in 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) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return running_loss / total, correct / total # 先跑 10 个 epoch 做基线 for epoch in range(10): loss, acc train_one_epoch(model, train_loader, criterion, optimizer, device) print(fepoch {epoch1}: loss{loss:.4f}, acc{acc:.4f})注意这里没有写验证逻辑先用训练集准确率做检查。如果训练集上损失根本不下降问题大概率出在数据加载或者标签映射上不是模型问题。基线跑通后再补验证集评估。全连接层替换代码里model.fc.in_features是 PyTorch 提供的接口不需要自己写死 512这样换主干网络时不用改这行代码。CrossEntropyLoss内部直接接收类别索引不需要对标签做 one-hot 编码。4. 数据增强与小样本调参让 94 张训练图发挥出 200 张的效果模型能跑通只是第一步。3400 张数据要训练出在真实场景里能用的模型必须把数据增强做到位否则验证集准确率会很难看。这一章讲的是我在参数层面的核心取舍也是这份数据集上差距最大的部分。4.1 增强策略按“单类 94 张”设计小样本分类的增强原则是训练集越“折腾”模型越稳但折腾幅度过大会改变图像本身的类别特征。果蔬分类里最典型的例子就是颜色——番茄的红色、香蕉的黄色是核心判别特征如果ColorJitter的亮度、饱和度幅度调太高红色变成橙色、黄色变成灰黄色等于训练标签被污染了。from torchvision import transforms # 适合果蔬分类的增强组合 train_transform transforms.Compose([ transforms.RandomResizedCrop(size224, scale(0.6, 1.0), ratio(0.75, 1.333)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.02), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])RandomRotation(degrees15)控制旋转角度为 ±15 度。果蔬拍摄时不会有大幅度的倾斜旋转超过 30 度就会出现大量的黑色填充区域模型要花参数去学习“忽略黑边”这个无关模式。hue0.02是关键——色调偏移只给 0.02稍大一点就会让青椒变紫、番茄变橙。这里每个参数的取值逻辑是先往大了调跑两个 epoch 看验证集是否下降如果下降再往回缩。之前提过RandomResizedCrop的scale用(0.6, 1.0)这里补充说明ratio参数控制裁剪框的宽高比默认(0.75, 1.333)对圆形果实和长条形蔬菜都适用。如果数据集里叶菜类比如菠菜、生菜比例高可以把ratio放宽到(0.5, 2.0)模拟叶片被拉长变形的状态能小幅提升这部分的泛化能力。4.2 类别不均衡与标签平滑统计完每类样本数后如果发现某些类别明显少于平均值训练时要让损失函数对少数类更敏感。常见做法是给CrossEntropyLoss传入weight参数权重取每类样本数的倒数。计算方式如下import torch import numpy as np # 统计每类样本数假设 train_dataset.targets 是类别索引列表 targets train_dataset.targets class_counts np.bincount(targets, minlength36) # 用样本数倒数并归一化避免个别类别权重过大导致训练震荡 weights 1.0 / class_counts.astype(np.float32) weights weights / weights.sum() * len(weights) criterion torch.nn.CrossEntropyLoss(weighttorch.FloatTensor(weights).to(device))按倒数加权是实践中最好用的一档。比它强烈的有WeightedRandomSampler——通过过采样让每个 batch 里都出现足够多少数类样本但对 3400 张的数据集来说少数类反复被采样容易让模型记住特定样本的噪声。比它温和的有(1 - beta) / (1 - beta^count)这类平滑公式适合极端不均衡的场景36 类果蔬远没到那个程度。另外建议在训练的前几个 epoch 用label_smoothing0.1。果蔬分类存在大量视觉相似对比如青苹果和青梨、多种辣椒之间硬标签会把置信度压向 1.0模型在低层特征上过拟合。标签平滑让正确类别的目标概率变成 0.9剩下 0.1 均分给其他 35 类这类相似对的泛化通常会好 1 到 2 个百分点。4.3 训练超参数推荐小样本迁移学习有几个先手参数表格里列的是我在类似项目里的稳定起点参数推荐值说明batch size323400 张图batch 太大导致每 epoch 更新次数太少初始学习率1e-4冻结主干时可以用 1e-3解冻后必须降到 1e-4优化器AdamW配合 weight_decay1e-3抑制 FC 层过拟合epoch30~50表现最好的 checkpoint 通常在 15~25 epoch 之间学习率调度CosineAnnealingLRT_max30比阶梯下降更稳冻结策略先冻结主干跑 10 epoch再解冻全量微调避免随机初始化的 FC 层拖坏主干梯度最后一个“冻结策略”值得展开。替换后的全连接层是随机初始化如果一开始就全量微调前向传播的梯度很大会把预训练主干的特征直接冲乱。先冻结主干只训练 FC 层让输出头先适应 36 类分布10 到 15 个 epoch 后主干特征已经在正常的特征空间里了再解冻做整体微调学习率降到 1e-4稳定性好很多。5. 用混淆矩阵和 ONNX 导出验证模型质量训练结束后的第一步不是看准确率数字而是看模型错在哪里。36 类果蔬里必然存在系统性的混淆对——比如圆茄子和紫洋葱、青柠和酸橙准确率看不出来但混淆矩阵一眼就能暴露。这也是我判断模型是否真正可交付的最后一道工序。5.1 混淆矩阵定位易混类别对import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix import numpy as np # model 切换到 eval 模式跑完整个验证集收集预测结果 model.eval() all_preds [] all_labels [] 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) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) # 找出混淆最多的 top-5 类别对跳过对角线 np.fill_diagonal(cm, 0) confused_pairs [] for i in range(len(cm)): for j in range(len(cm)): if cm[i][j] 0: confused_pairs.append((cm[i][j], i, j)) confused_pairs.sort(reverseTrue) for count, i, j in confused_pairs[:5]: class_i val_dataset.classes[i] class_j val_dataset.classes[j] print(f{class_i} 被识别为 {class_j}: {count} 张)这段代码输出的不是 png 图片而是直接打印混淆最严重的 5 对类别名。np.fill_diagonal(cm, 0)把对角线清零否则对角线上全是正确分类的数排序后最前面的一定是自己对自己看不到真正的问题对。打印出的类别对要重点检查如果某两类被混淆的次数超过该类验证样本的 30%说明这两类在视觉特征上确实接近更适合的做法是在标注阶段合并成一个大类而不是继续堆模型容量。5.2 导出为 ONNX最后的工程化一步验证矩阵确认无误后把 PyTorch 模型导出为 ONNX 格式便于脱离训练框架部署。导出时的关键点是固定输入尺寸否则动态尺寸会导致 ONNX Runtime 里的推理变慢。import torch.onnx model.eval() dummy_input torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, fruit_veg_cls.onnx, input_names[pixel_values], output_names[logits], dynamic_axes{pixel_values: {0: batch_size}, logits: {0: batch_size}}, opset_version17 )dynamic_axes里只允许 batch 维度动态图像尺寸维度保持固定的 224x224。导出的 ONNX 文件可以先用onnxruntime做一个快速的推理对齐测试用同一张验证图片分别跑 PyTorch 模型和 ONNX 模型对比输出logits是否一致误差在 1e-4 量级内算正常。这一道验证通过后模型就可以交给服务端了。剩下的工作无非是把 36 类索引映射表随模型一起发布出去并约定调用方传入的图像按Resize(256) CenterCrop(224)预处理与训练时的验证集保持一致。本文还有配套的精品资源点击获取
返回列表