尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

基于PyTorch和MobileNetV3的中草药图像识别实战解析

基于PyTorch和MobileNetV3的中草药图像识别实战解析 简介面向深度学习图像识别与中草药智能化鉴定场景这是一份完整可运行的毕业设计级项目源码包适合计算机专业学生、PyTorch初学者及对细粒度图像分类感兴趣的开发者。项目采用ResNet结合CBAM注意力机制配套10类共2500张中草药图像数据集并应用数据增强提升泛化能力。资源共2000个文件主体为1976张jpg样本图片另有12个py训练/推理脚本、训练日志、HTML可视化报告、配置文件及说明文档压缩包约198.37MB目录结构便于复现实验。目前已有203人学习下载。通过源码可学习模型搭建、超参数调优、训练评估全流程还可直接基于数据增强与注意力模块改造自己的分类任务为毕业设计或科研实践提供扎实参考。 我不是第一次碰中草药识别这类项目了但每一次重新做都会被数据问题磨掉一层皮。这次梳理的是一个基于深度学习的端到端中草药图像识别项目代码已经整理好放在开源仓库里。项目本身不算复杂用CNN做图像分类主流的PyTorch框架MobileNetV3作为backbone训练集来自采集和公开数据混合覆盖了大概30种常见中草药。整套流程跑下来单卡GPU训练一个下午就能出可用模型在验证集上Top-1准确率能到85%以上。如果你正准备接触深度学习图像分类或者手头恰好有中草药识别的需求这个项目源码的参考价值会非常大。我先把结论摆在这里中草药识别这个场景难的不是模型是数据。你去看网上很多类似的论文和项目模型结构都差不多ResNet也好、EfficientNet也好拉开差距的地方全在训练数据的数量、质量和分布设计上。而数据问题在通用图像分类里往往是用来“一带而过”的放到中草药识别里却会被无限放大——因为中草药本身太特殊了。后面我会详细拆先讲整体设计。1. 项目整体设计与技术选型1.1 为什么用深度学习而不是传统图像识别很多没接触过图像识别的人会问一个问题中草药识别用传统图像处理不行吗比如采集叶片形状、纹理特征提取颜色直方图再利用SVM、随机森林这类传统机器学习算法做分类。坦白说如果只做3到5种差异极大的草药传统方法确实能跑通。但一旦类别数上升到20种以上形态相近的草药扎堆出现传统手工特征的区分能力就完全不够了。拿最常见的例子来说薄荷和荆芥的叶片形状都是卵圆形颜色都是绿色边缘都有锯齿光靠形状、纹理这些手工特征别说机器人眼都有点犯迷糊。但深度学习模型通过多层卷积自动学习到的特征可以捕捉到人眼不容易描述的细微差异——叶脉走向、表面绒毛密度、叶缘锯齿的深浅规律等。这就是为什么这个项目直接选择了深度学习路线而不是去走传统特征的弯路。1.2 网络结构选型的思路为什么是CNN而不是Transformer当前深度学习做图像分类两大主流方向是CNN卷积神经网络和Vision TransformerViT。按理说ViT在大型数据集上表现极其出色为什么这个项目还是选用了CNN核心原因是三个字数据量。ViT是一种极度“吃数据”的模型结构它不像CNN那样本身就带有“局部性”和“平移不变性”的归纳偏置所以需要海量数据才能学到像样的特征表达。在ImageNet那种千万级数据集上ViT固然很强但在这个项目中草药数据集只有一两万张图像ViT的网络优势不仅发挥不出来反而会因为数据不足引发严重的过拟合——训练集准确率接近100%验证集却掉到60%以下。反观CNN尤其是轻量级结构MobileNetV3参数量小、归纳偏置强、在中小规模数据上表现稳定还非常适合后期部署到移动端做田间地头的实时识别。所以这个项目最终选择了MobileNetV3-Large作为主干网络。1.3 开发框架与训练环境选型PyTorch是这个项目的首选框架没有悬念。在当前深度学习开源生态里PyTorch在学术界和工业界的覆盖率已经非常高了。从源码的易读性、Debug的直观性动态图机制到TorchVision中预训练模型的下载便利性PyTorch都做得非常成熟。测试下来同等熟练度下用PyTorch写一个分类训练脚本代码量比TensorFlow要少大概三分之一而且不用被静态图的各种shape声明折磨。训练环境方面建议Windows或Linux系统配一张NVIDIA显卡哪怕入门级的GTX 1660都行CUDA版本需要提前匹配好。没有GPU的话纯CPU也能跑但训练时间会从几十分钟撑到十几个小时基本没法做多轮实验调参。我在环境配置上踩过一次大坑后面专门用一节来讲。2. 数据集准备与预处理2.1 中草药图像数据的特殊难点这是整个项目最核心、最需要重视的部分我花的时间占比超过70%。中草药图像数据有三个非常特殊的问题第一类间相似度高。刚才说过很多不同种类的草药外观极其接近有的甚至只在叶片背面的绒毛密度上有差异。对分类模型来说这等于是在做“找不同”游戏非常考验特征的感知能力。第二类内差异巨大。同一种草药幼苗期和成熟期可能长得完全不同干燥药材和新鲜植株也完全是两个样子。如果训练数据只覆盖了其中一种状态模型的泛化能力就会严重受挫。比如金银花新鲜的时候是黄白相间的花朵干燥入药后是扭曲的暗黄色条状物两个形态差异大到仿佛是不同的物种。第三背景干扰严重。真实场景下拍摄的中草药图像背景里往往包含泥土、杂草、其他植物、甚至手指和手机支架。模型如果只见过纯色背景下的标准图换到野外背景就会直接“翻车”。2.2 数据采集的三种途径这个项目的数据集由三部分拼合而成我自己相机实地拍摄的样本、中草药植物园收集的图像、以及网上开源数据集和搜索引擎里筛选出来的公开图片。三种来源比例大致是3:4:3。实地拍摄时我总结出一套采集规范现在整理给各位提示每类药材至少采集300张基础图像每张图像尽量使药材主体占据画面50%以上同时刻意保留部分背景元素。除了正面平视还要采集俯视、侧视、逆光等不同角度。同一株药材要分别拍嫩叶期、成熟期和花朵/果实期。公开图像筛选时要特别小心“错标”问题。搜索“蒲公英”时很容易混入苦苣菜、续断菊这类外形相似的植物。我采取的策略是先人工粗筛一遍拿不准的图直接删除绝不抱着“先留一张反正影响不大”的心态——分类任务中一个错误的标注可能误导整个类别的特征学习。2.3 数据标注与格式统一标注工作使用LabelImg完成这个是业界比较通用的图像标注工具。不过这个项目做的是图像分类而非目标检测所以不需要画边界框只需要按照类别名称建立文件夹结构即可。最终的数据存储格式是dataset/ ├── train/ │ ├── 薄荷/ │ │ ├── mint_001.jpg │ │ ├── mint_002.jpg │ │ └── ... │ ├── 金银花/ │ ├── 野菊花/ │ └── ... └── val/ ├── 薄荷/ ├── 金银花/ └── ...训练集和验证集按照8:2比例随机划分。这里要特别提醒划分时必须基于类别而非整图混分确保每个类别在训练集和验证集中都有足够且均衡的代表。2.4 数据增强策略小数据集救星数据增强是这个小数据集项目能跑出来的关键所在。简单理解数据增强就是“用有限的数据变出更多的数据”。我在项目中采用了如下的增强组合随机水平翻转概率0.5让模型对镜像不敏感因为拍摄时左右方向不固定随机旋转±15度模拟拍摄角度的小幅变化随机缩放裁剪范围0.81.0模拟远近不同的拍摄距离颜色抖动亮度±20%、对比度±15%、饱和度±10%模拟不同光照条件随机擦除概率0.3模拟叶片被遮挡的工程场景这一套增强打下来每个epoch模型看到的图像都略有不同等于把13000多张训练图“变”成了几乎无限多。实验对比显示加增强和不加增强的验证集准确率差距高达8到10个百分点可见常规训练中这一步不能省。3. 模型实现与训练调参3.1 核心模型搭建迁移学习策略全天下做图像分类的工程师100个人里有95个人会用迁移学习这个项目也不例外。所谓迁移学习通俗讲就是“站在巨人的肩膀上”——利用一个已经在ImageNet1400万张图像的庞大数据库上训练好的模型作为起点把前面卷积层学会的通用特征提取能力直接借来用只替换掉最后面的全连接分类层让它适应“中草药分类”这个新任务。源码中模型搭建的核心代码是这样的import torch import torch.nn as nn from torchvision import models def build_model(num_classes30, pretrainedTrue): model models.mobilenet_v3_large(pretrainedpretrained) # 获取原模型最后分类层的输入维度 in_features model.classifier[-1].in_features # 替换新的分类器适配当前任务 model.classifier[-1] nn.Linear(in_features, num_classes) return model模型一旦换了新的分类层整个网络就有了“混搭”结构前几层卷积参数用ImageNet预训练好的初始值最后一层线性层是随机初始化的。训练时如果从头到尾都用同样的学习率相当于让“新同学”和“老司机”迈同样的步子——前面已经学好的特征很可能被大幅破坏。我采用的策略是切分学习率backbone部分的初始学习率设为0.0001而新加的分类层用0.001。这样一来预训练特征只做微调新分类层则大步快跑地快速收敛。3.2 损失函数与优化器选择分类问题最经典的损失函数就是交叉熵损失CrossEntropyLoss没有花里胡哨的必要。它做的事情可以用大白话来描述如果模型对正确类别的预测概率是0.9那么损失就是-log(0.9)≈0.105如果概率只有0.1损失就涨到-log(0.1)≈2.3。损失越小说明模型预测得越准。优化器方面我选用AdamW在Adam的基础上引入了权重衰减的修正能在保证快速收敛的同时有效抑制过拟合。初始学习率0.001配合CosineAnnealingLR余弦退火调度器让学习率在整个训练过程中从0.001平滑地降到接近于0相当于前期大步探索后期小步精修。3.3 训练轮数与精度的关系一个关键的曲线观察很多刚入门的人都会问“该训练多少个epoch”。我的回答永远都是别拍脑袋定看曲线说话。这次项目的实验中我记录了每一轮的训练准确率和验证准确率表格如下训练轮数训练集损失训练集准确率验证集准确率51.78232.5%28.3%100.93663.1%55.7%150.54781.2%71.3%200.31290.4%78.6%250.18495.8%83.1%300.09698.7%85.2%350.05199.6%84.9%400.02899.9%84.3%观察这个表格可以看到非常典型的规律前20轮训练集和验证集的准确率同步上升模型确实在学习有效特征但从30轮之后训练集准确率还在缓慢爬升98.7%→99.9%验证集准确率反而开始下降了85.2%→84.3%。这就是标准的过拟合信号——模型开始“死记硬背”训练图像中的细节和噪声而不是提取泛化性强的普适特征。所以这个项目最终选定的训练轮数是30对应一个早停策略Early Stopping当验证准确率连续5轮不再提升就提前终止训练并回滚到最佳模型参数。3.4 类别不均衡问题的处理采集数据时有一种很现实的情况像蒲公英、车前草这类分布极广、随处能拍到的药材轻轻松松收集了1000多张图而像雪莲花、铁皮石斛这类比较稀有、地域性强的药材勤勤恳恳一个月也攒不到100张。这种类别样本数量差异较大的情况在深度学习里会引发一个严重的偏向问题模型会把大量精力花在样本多的类别上对样本少的类别直接“摆烂”。举个例子如果70%的训练数据都是蒲公英模型只要把所有图都预测为蒲公英就已经拿到70%的准确率了它没必要费力去学其他类别。处理这个问题我用了两板斧第一对样本少的类别提高采样权重——在DataLoader中使用WeightedRandomSampler让每个类别在每个epoch采样的概率接近均衡第二对样本少的类别做更强的数据增强如把颜色抖动幅度加大、随机擦除概率提高相当于“人工多给它一些变体”。这套组合下来原本稀有的铁皮石斛识别准确率从54%提升到了78%效果还是相当显著的。4. 完整训练与评估流程4.1 数据准备与预处理入口源码里数据读取部分的实现逻辑是这样的from torch.utils.data import DataLoader from torchvision import datasets, transforms train_transform transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.RandomResizedCrop(size224, scale(0.8, 1.0)), transforms.ColorJitter(brightness0.2, contrast0.15, saturation0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_dataset datasets.ImageFolder( rootdataset/train, transformtrain_transform ) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4 )这里的Normalize操作用的mean和std是ImageNet数据集的统计值。因为模型是用ImageNet预训练的输入数据也应该服从相同的分布。如果图省事不归一化或者用错统计值微调效果就会打折扣。很多新手容易在这里踩坑我见过不少项目模型训不动最后发现是数据预处理没对齐。4.2 训练主循环源码解读训练核心代码逻辑不复杂核心流程就是“正向计算损失→反向传播梯度→优化器更新参数”。但有几个细节值得强调for epoch in range(epochs): model.train() 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() # 每轮结束后做验证 val_acc evaluate(model, val_loader, device) print(fEpoch {epoch1}/{epochs}, Loss: {loss.item():.4f}, Val Acc: {val_acc:.2f}%) # 保存最佳模型 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth)第一个细节是optimizer.zero_grad()必须在每个batch开始前调用。PyTorch的gradient默认是累积的不清零的话多个batch的梯度就会叠加在一起导致参数更新方向严重偏离。这是我见过新手最容易犯的错之一。第二个细节是model.train()和model.eval()的切换。因为模型中包含BatchNorm层和Dropout层训练和推理时的行为不一样。训练时使用批量统计信息推理时使用全局统计信息忘记切换的话验证结果会异常地差。第三个细节是最佳模型的保存策略。训练过程中我只保存验证集上表现最好的那一份权重best_model.pth而不是最后一次epoch的权重。因为最后一轮往往已经过拟合了而验证集准确率最高的那个历史节点才是泛化能力最强的。4.3 评估指标与混淆矩阵分析最终在30类中草药的验证集上模型整体Top-1准确率达到了85.2%Top-5准确率即答案包含在模型预测概率最高的前5个类中达到了96.7%。作为参考医疗领域的植物识别如果Top-5能到95%以上已经具备辅助实际使用的条件了。但只看整体准确率远远不够。我额外绘制了混淆矩阵Confusion Matrix来观察哪两类药材最容易被混淆。结果毫不意外混淆程度最高的是薄荷和荆芥这个经典的近似类有约9%的薄荷图像被误判为荆芥其次是干燥的白芍和白芷因为两者经过干燥处理后的颜色、纹理非常接近。这类易混淆问题如果还想进一步优化可以采用的方案包括采集更多这类样本的精细图像、在训练时对易混淆类别的样本做更密集的采样或者干脆设计一个两级分类结构——先把大差异的类别分开再由第二个模型专门区分易混淆类别。考虑到项目规模和投入产出比当前阶段暂时不采用这么重的方案但如果在工业界实际落地这是很清晰的一条后续迭代路径。5. 常见问题与排查技巧实录5.1 训练Loss不下降怎么办这是所有人都会遇到的场景模型结构对着论文写对了数据加载没问题训练脚本跑起来了但loss像是焊死在初始值附近5轮10轮过去了一点动静没有。根据这个项目的排查经验优先级最高的两个检查点是第一标签和模型输出维度是否对齐。如果数据集有30个类别但模型最后的线性层输出维度设成了1000忘了改loss就会一直在高位震荡。这个问题非常隐蔽因为代码不会报错就是loss不降。验证方法很简单打印一个batch的logits和label的shape逐项核对。第二学习率是否过小或过大。学习率过小会导致模型的收敛速度极慢感观上就像“没在训练”。反之学习率过大则可能导致loss不降反升。建议在调参初期使用学习率预热Learning Rate Warmup和快速衰减实验先用0.001训练20个batch观察loss变化趋势如果没有明显下降再尝试0.01或0.0001。5.2 验证集准确率高但实际识别效果差怎么回事训练完成后把模型部署到手机上拍一张真实的草药照片结果识别结果完全不对。这个问题本质上是一个典型的“数据集偏移”Dataset Shift问题。详细说起来训练数据里很大一部分是理想的拍摄条件光线均匀、背景纯净、药材居于画面正中央。但在真实使用场景中用户拍摄的照片可能是阴天的弱光、杂乱的草丛背景、药材只占画面一角、甚至还有手指遮挡。解决方案有两个方向。第一个是训练阶段做更“狠”的数据增强我补上了随机擦除、马赛克增强和背景混合等策略让模型在训练期就见过各种“脏乱差”的输入。第二个是推理阶段加预处理在模型真正执行分类前先运行一个轻量级目标检测模型把画面中的药材区域裁剪出来再送入分类网络。后者虽然技术上更复杂但工程效果显著更好。提示如果觉得部署额外的检测模型太重可以退而求其次——输入图先做中心裁剪Center Crop强制模型把注意力集中在图像中央区域。中草药识别的实际场景中用户大概率会把药材放在画面中心再拍照中心裁剪能够排除大部分背景干扰。5.3 训练时OOM显存不足的排查思路这个项目在batch_size64训练时我在一张8G显存的显卡上直接OOM了。降低图片分辨率或者减小batch_size都可以快速缓解。但这里有一个技巧与其盲目降低batch_size不如先检查是不是开启了过多的DataLoader工作进程。num_workers开得太大时虽然不会直接占用GPU显存但可能引发内存交换风暴间接拖慢训练效率甚至导致进程被系统杀掉。最终我的稳定配置是batch_size32、图片分辨率224x224、混合精度训练AMP。混合精度训练是一个性价比极高的优化选项——在大多数NVIDIA显卡上开启AMP后显存占用能下降40%左右同时训练速度还能提升30%左右而且模型精度基本不受影响。PyTorch自带的torch.cuda.amp接口用起来很方便几行代码就能接入。5.4 Windows环境下深度学习环境配置的几个坑之前提到我在环境配置上踩过大坑这里展开细说。PyTorch安装本身不复杂复杂的是环境依赖的版本匹配问题。CUDA Toolkit、显卡驱动、PyTorch三者的版本必须兼容否则会出现装好了import就报错的情况。我在Windows上实测的稳定组合是NVIDIA驱动版本535及以上、CUDA 11.8、PyTorch 2.0cu118。先安装显卡驱动再安装CUDA Toolkit最后用pip安装对应版本的torch这个顺序不要乱。另一个常见坑是虚拟环境。强烈建议用conda创建独立的Python 3.9环境而不是直接使用系统Python。深度学习依赖的包版本冲突非常频繁今天装的某个包升级了明天就可能把torch的某个依赖顶掉。独立环境等于给自己建了一个隔离区怎么折腾都不怕。最后给一句真心建议有条件的话直接使用云GPU平台做训练比如AutoDL、恒源云这类按小时计费的共享GPU平台一度电费级别的成本就可以用上RTX 3090甚至A100。尤其适合前期快速验证模型可行性的阶段真的没必要为了跑一次完整训练特意去买一块昂贵显卡。6. 源码核心模块与实际使用说明6.1 源码整体结构与运行流程项目源码的目录结构非常清晰拿到手之后不用读文档就能猜到大概herbal-recognition/ ├── checkpoints/ # 模型权重存放目录 ├── dataset/ # 数据集目录 ├── models/ │ ├── __init__.py │ └── mobilenet.py # 模型结构定义 ├── utils/ │ ├── data_loader.py # 数据加载与增强 │ ├── train.py # 训练逻辑 │ ├── evaluate.py # 评估逻辑 │ └── inference.py # 单张图片推理 ├── config.py # 全局配置 ├── train.py # 训练入口 └── predict.py # 预测入口运行流程非常直接。训练阶段修改config.py中数据路径和超参数然后运行python train.py推理阶段运行python predict.py --image test.jpg --checkpoint checkpoints/best_model.pth控制台上就会输出Top-5预测结果以及对应的置信度。6.2 推理代码从加载权重到输出预测结果推理部分的代码写得比较直白一行一行看很容易理解import torch from PIL import Image from torchvision import transforms from models.mobilenet import build_model def predict(image_path, checkpoint_path, class_names): # 加载模型并切换到推理模式 model build_model(num_classeslen(class_names), pretrainedFalse) model.load_state_dict(torch.load(checkpoint_path, map_locationcpu)) model.eval() # 预处理与模型推理 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0) with torch.no_grad(): outputs model(input_tensor) probs torch.softmax(outputs, dim1) top5_probs, top5_indices torch.topk(probs, k5) for i in range(5): idx top5_indices[0][i].item() print(fTop {i1}: {class_names[idx]} ({top5_probs[0][i].item()*100:.2f}%))有个细节要注意单张推理和批量训练时的图片预处理要完全一致连Resize的方式都要统一。比如训练时用了RandomResizedCrop推理时就不能只是简单的Resize否则输入分布不一致预测结果有偏差。另外推理阶段最好加上with torch.no_grad()关闭梯度计算能显著降低显存占用和加速计算。6.3 后续优化方向部署与模型压缩如果这个项目最终要落地成一个小程序或App应用模型的大小和推理速度就变成核心指标了。MobileNetV3的权重文件大约20MB这个体积基本可以接受但还有进一步压缩的空间量化Quantization是一种常见手段把参数从32位浮点数压缩到8位整数模型体积直接缩小到四分之一推理速度却成倍提升精度损失通常只有1%到2%。另外一个可行的方向是模型蒸馏训练一个参数量更大的教师模型比如ResNet50获得更高的精度上限再用这个教师模型的软标签去指导MobileNetV3学生模型的学习。这样可以在不增加推理代价的前提下把MobileNetV3的准确率再提升两三个百分点。我在后续版本迭代中已经在实验这个方向测试结果出来后有机会再单独整理一篇经验分享。7. 写在最后的实操心得这个项目从零开始到跑通完整训练流程最深的体会是深度学习项目里真正值钱的部分是数据和工程细节模型结构反而是固定的“标准件”。数据是否干净、增强是否到位、超参是否合理、训练还是推理的模式切换有没有做到位每一步都在影响最终效果。很多人在搭建好模型后急着开启训练结果准确率不高又不知道是哪一环出了问题于是反复堆数据、换网络——问题的根源往往不在那里。最后再分享一个实战小技巧保存模型权重文件时建议顺便保存一份训练日志每一轮的loss和准确率曲线数据格式可以是CSV或者JSON。后续回看项目时这份日志能让你快速复盘当时的训练过程还能用来自动生成训练曲线图。有这份记录在手做起实验对比来会轻松很多。本文还有配套的精品资源点击获取
返回列表