☰
基于PyTorch和MobileNetV3的中草药图像识别实战解析
2026/10/9 15:52:02 网站建设 项目流程

简介:面向深度学习图像识别与中草药智能化鉴定场景,这是一份完整可运行的毕业设计级项目源码包,适合计算机专业学生、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 Transformer(ViT)。按理说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.8~1.0):模拟远近不同的拍摄距离
  • 颜色抖动(亮度±20%、对比度±15%、饱和度±10%):模拟不同光照条件
  • 随机擦除(概率0.3):模拟叶片被遮挡的工程场景

这一套增强打下来,每个epoch模型看到的图像都略有不同,等于把13000多张训练图“变”成了几乎无限多。实验对比显示,加增强和不加增强的验证集准确率差距高达8到10个百分点,可见常规训练中这一步不能省。

3. 模型实现与训练调参

3.1 核心模型搭建:迁移学习策略

全天下做图像分类的工程师,100个人里有95个人会用迁移学习,这个项目也不例外。所谓迁移学习,通俗讲就是“站在巨人的肩膀上”——利用一个已经在ImageNet(1400万张图像的庞大数据库)上训练好的模型作为起点,把前面卷积层学会的通用特征提取能力直接借来用,只替换掉最后面的全连接分类层,让它适应“中草药分类”这个新任务。

源码中模型搭建的核心代码是这样的:

import torch import torch.nn as nn from torchvision import models def build_model(num_classes=30, pretrained=True): model = models.mobilenet_v3_large(pretrained=pretrained) # 获取原模型最后分类层的输入维度 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(size=224, scale=(0.8, 1.0)), transforms.ColorJitter(brightness=0.2, contrast=0.15, saturation=0.1), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) train_dataset = datasets.ImageFolder( root='dataset/train', transform=train_transform ) train_loader = DataLoader( train_dataset, batch_size=32, shuffle=True, num_workers=4 )

这里的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(f"Epoch {epoch+1}/{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_size=64训练时,我在一张8G显存的显卡上直接OOM了。降低图片分辨率或者减小batch_size都可以快速缓解。但这里有一个技巧:与其盲目降低batch_size,不如先检查是不是开启了过多的DataLoader工作进程。num_workers开得太大时,虽然不会直接占用GPU显存,但可能引发内存交换风暴,间接拖慢训练效率甚至导致进程被系统杀掉。

最终我的稳定配置是:batch_size=32、图片分辨率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.0+cu118。先安装显卡驱动,再安装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_classes=len(class_names), pretrained=False) model.load_state_dict(torch.load(checkpoint_path, map_location='cpu')) 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, dim=1) top5_probs, top5_indices = torch.topk(probs, k=5) for i in range(5): idx = top5_indices[0][i].item() print(f"Top {i+1}: {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。后续回看项目时,这份日志能让你快速复盘当时的训练过程,还能用来自动生成训练曲线图。有这份记录在手,做起实验对比来会轻松很多。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询