尧图精选

中草药叶片识别实战:数据划分、可视化与PyTorch迁移学习全流程

🕒 发布时间:2026/10/1 3:50:30 📁 来源:尧图网络
简介面向深度学习与图像分类的中草药叶片识别数据集适合希望快速获取已标注图像数据以验证模型效果的开发者与研究人员。数据已按训练集与验证集划分完毕训练集含4800张图片、验证集含1400张图片覆盖80个中草药叶片类别可直接作为YOLOv5分类任务的训练数据免去自行采集、清洗与整理的耗时环节。压缩包共2000个文件主体为1998张JPG叶片图片另附1个JSON类别字典文件提供80个类别名称与编号的完整对应关系。另有1个Python数据可视化脚本可随机抽取4张图片进行展示并将结果保存到当前目录便于快速核查数据质量与分布情况。包体约175MB目录结构清晰训练集与验证集分目录存放开箱即可用目前已有119人下载学习可为中草药识别相关的模型训练、算法对比与论文实验提供直接可用的数据支撑。1. 中草药叶片识别数据没组织好再好的模型也白搭做中草药叶片识别分类时大多数人第一步是去下载预训练模型结果换了好几个网络准确率就是上不去。我见过太多人卡在这个环节最后才发现问题根本不在模型而在最前面那几步——数据怎么划分、类别字典怎么生成、图像有没有真的看一遍。标题里给的划分好的数据、类别字典文件、python数据可视化脚本恰恰是这类计算机视觉小项目最容易忽视的地基。这篇文章不聊高深理论直接讲我怎么组织一份中草药叶片数据集从目录结构到训练验证把每一步代码和参数讲透。适合正在做课程大作业、竞赛或药植园课题的从业者尤其是被乱糟糟的图片文件逼疯的那批人。2. 数据集怎么才算划分好目录约定与类别字典生成2.1 从原始照片到标准目录train / val / test 应该怎么摆很多人的原始素材是这样一个文件夹里面按类别放了子文件夹比如银杏叶/、薄荷叶/、蒲公英/每类几十张照片。这是好事至少类别信息有了。但直接拿它训练你会踩大坑没有验证集和测试集调参全靠感觉最后汇报的准确率可能是模型在训练集上背出来的。最常见的做法是整理成 ImageNet 风格的目录dataset/ train/ 银杏叶/ 001.jpg 002.jpg ... val/ 银杏叶/ ... test/ 银杏叶/ ...train 用于更新模型参数val 用于每轮训练后评估、决定要不要调整学习率或早停test 只在全部训练结束时跑一次模拟真实使用场景。三者比例我一般用 7:2:1。如果你的原始数据不多比如每类只有 30 张那就把 val 缩到 15%test 保持 15%train 70%——宁可 train 少一点也要保证 val 和 test 是没见过的个体。如果你连按类别分文件夹都没有那就得先人工整理。看到这里你可能觉得麻烦但相信我这一步省掉的每一分钟之后都会在调试模型时加倍还回来。2.2 类别字典文件为什么需要它格式怎么定模型输出的不是一个汉字类别名而是一个整数索引。你训练时可能会把银杏叶映射到 0薄荷叶映射到 1但问题在于这个映射关系如果不落盘保存到推理阶段就全乱了。类别字典文件通常是一个 JSON就是用来固化这个映射的它只有两行内容但缺了它模型就是个黑匣子。下面这个脚本扫描目录并生成class_indices.jsonimport os import json from pathlib import Path data_root Path(dataset/train) # 按名称排序保证在 Linux 和 Windows 上顺序一致 classes sorted([p.name for p in data_root.iterdir() if p.is_dir()]) class_to_idx {cls: i for i, cls in enumerate(classes)} idx_to_class {str(i): cls for cls, i in class_to_idx.items()} with open(class_indices.json, w, encodingutf-8) as f: json.dump({class_to_idx: class_to_idx, idx_to_class: idx_to_class}, f, ensure_asciiFalse, indent2) print(类别数量:, len(classes)) print(映射表:, class_to_idx)classes sorted(...)这行非常关键。不同操作系统上os.listdir的返回顺序可能不一样如果不排序同一批图片在 Windows 上训练、到 Linux 上推理索引就会错位。我一般只保留idx_to_class这份映射作为字典文件因为推理时拿到的是索引你需要查回名字训练时则反向使用。ensure_asciiFalse让中文直接写进 JSON而不是变成\u67ab...肉眼可读方便你后来核验。2.3 自动划分脚本与随机种子让一次划分可靠复用手动去拖几百张图片到三个文件夹是体力活而且容易出错。写一个脚本自动划分顺便固定随机种子这样下次换一批数据、或者别人复现你的实验时结果是一致的。import shutil import random from pathlib import Path src Path(原始叶片数据) # 里面是按类别分的子文件夹 dst Path(dataset) # 目标根目录 random.seed(42) # 固定随机种子保证可复现 train_ratio, val_ratio 0.7, 0.2 for cls_dir in src.iterdir(): if not cls_dir.is_dir(): continue images list(cls_dir.glob(*.jpg)) list(cls_dir.glob(*.png)) random.shuffle(images) n_train int(len(images) * train_ratio) n_val int(len(images) * val_ratio) for phase, subset in [ (train, images[:n_train]), (val, images[n_train:n_train n_val]), (test, images[n_train n_val:]), ]: out_dir dst / phase / cls_dir.name out_dir.mkdir(parentsTrue, exist_okTrue) for img in subset: # 用 copy 而不是 move保留原始数据磁盘紧张时再改 move shutil.copy(img, out_dir / img.name) print(划分完成)注意我用shutil.copy而不是move。为什么因为原始照片是不可再生的万一划分逻辑有 bugmove 之后数据就乱了copy 顶多多占点磁盘。每类图片数量不一致时n_train和n_val按各自比例截断我这些年见过的坑是有人用固定数字n_train 50去切遇到只有 20 张的类别切出来 train 为空模型直接报错。所以记住划分永远按比例不按固定数量。3. 数据可视化脚本训练前先看清楚你的数据3.1 用 matplotlib 看类别分布发现数据不平衡很多中草药叶片数据集天然不平衡——银杏叶好采可能拍了 500 张某种珍稀药用植物只能找到 30 张。如果不看这个分布训练出的模型会严重偏向样本多的类别少数类一张都不认得。我习惯在训练前先跑一个统计脚本import matplotlib.pyplot as plt from pathlib import Path import collections data_root Path(dataset/train) counts collections.Counter() for cls in sorted([p.name for p in data_root.iterdir() if p.is_dir()]): counts[cls] len(list((data_root / cls).glob(*.jpg))) counts[cls] len(list((data_root / cls).glob(*.png))) plt.figure(figsize(10, 6)) plt.bar(counts.keys(), counts.values()) plt.xticks(rotation45, haright) plt.ylabel(样本数量) plt.title(train 集类别分布) plt.tight_layout() plt.savefig(class_distribution.png, dpi150)这个脚本输出的class_distribution.png会直接告诉你该不该做类别均衡采样。如果某个类别的数量只有最大类的 1/10训练时就要用WeightedRandomSampler或者做数据增强迂回补齐而不是让模型自生自灭。另外我还会顺手打印min(counts.values())最小类数量少于 20 时再厉害的模型也救不回来你得想的是去补拍照片而不是换网络结构。3.2 随机抽样显示叶片样本检查图像质量与标签错误可视化不只是为了发论文凑图它是揪标签错误的最高效手段。一张图被贴错标签用肉眼可能几秒钟就能发现但模型会忠实地把错误学进去。下面这个脚本从每个类别里随机抽 3 张拼成一张网格图import matplotlib.pyplot as plt from PIL import Image import random from pathlib import Path data_root Path(dataset/train) classes sorted([p.name for p in data_root.iterdir() if p.is_dir()]) random.seed(42) samples_per_class 3 fig, axes plt.subplots(len(classes), samples_per_class, figsize(3 * samples_per_class, 3 * len(classes))) for i, cls in enumerate(classes): imgs list((data_root / cls).glob(*.jpg)) list((data_root / cls).glob(*.png)) random.shuffle(imgs) for j in range(samples_per_class): ax axes[i][j] if len(classes) 1 else axes[j] img Image.open(imgs[j]).convert(RGB) ax.imshow(img) ax.set_title(cls, fontsize9) ax.axis(off) plt.tight_layout() plt.savefig(sample_grid.png, dpi120)我强调convert(RGB)是因为有些相机可能输出带透明通道的 PNG或者灰度图统一转成 RGB 避免后续 PyTorch 加载时通道数报错。看到网格图后重点看两类问题一是图像有没有叶片本身的残缺、失焦二是是否有不属于这个类别的叶片混进来。如果发现问题直接去源目录里删除或重新归类这张图片别想着反正就一张不影响大局——十张错图就能让模型在对应类别上彻底跑偏。3.3 图像尺寸与亮度分布为后续预处理提供依据不是所有相机拍出来的叶片照片都是正方形的长宽比各异会让模型的感受野和池化层处理变得别扭。统计一下尺寸分布决定全局统一 resize 的目标尺寸这个步骤能省下人眼排查的时间。import numpy as np from PIL import Image from pathlib import Path data_root Path(dataset/train) widths, heights [], [] for img_path in data_root.rglob(*.jpg): img Image.open(img_path) w, h img.size widths.append(w) heights.append(h) print(宽度范围:, np.min(widths), -, np.max(widths)) print(高度范围:, np.min(heights), -, np.max(heights)) print(平均宽高比:, np.mean(np.array(widths) / np.array(heights)))如果平均宽高比在 1.0 左右那直接 resize 到 224×224 不会有太大畸变如果比值明显偏离 1比如大部分是 3:4 竖构图我一般会先做中心裁剪成正方形再 resize而不是直接拉伸。这个统计同样适用于筛选严重损坏的图——如果某张图的宽度是 0 或高度异常小多半是坏文件在加载时直接跳过并在日志里打条 warning 更稳妥。4. 从数据集到分类模型PyTorch 训练脚本与关键参数4.1 用 PyTorch 自定义 Dataset 加载划分好的数据有了标准目录结构和类别字典后写一个torch.utils.data.Dataset子类来加载图像。这里要注意__getitem__返回的不只是图像和标签还应该返回一个索引或是图像路径方便后面排查误分类样本。import torch from torch.utils.data import Dataset from PIL import Image import json from pathlib import Path class LeafDataset(Dataset): def __init__(self, root, transformNone): self.root Path(root) self.transform transform self.samples [] # (path, label, class_name) # 读取类别字典 with open(class_indices.json, r, encodingutf-8) as f: mapping json.load(f)[class_to_idx] self.class_to_idx mapping # 遍历 train/val/test 下的每个类别目录 for cls in sorted([p.name for p in self.root.iterdir() if p.is_dir()]): for img in (self.root / cls).rglob(*.jpg): self.samples.append((str(img), self.class_to_idx[cls], cls)) for img in (self.root / cls).rglob(*.png): self.samples.append((str(img), self.class_to_idx[cls], cls)) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label, cls_name self.samples[idx] image Image.open(path).convert(RGB) if self.transform: image self.transform(image) return image, label, path # path 用于后续误判溯源关键点在于返回path。很多人只返回(image, label)模型训练时没问题但出错了想追查是哪张图出错就只能靠猜。把路径带出来训练日志里每条 error 都能精确对应到某个文件这在做误判分析时价值极高。rglob(*.jpg)只能匹配一层子目录如果你的目录里还有嵌套目录建议把所有图片扁平到一个类别文件夹下避免漏加载。4.2 迁移学习ResNet18 微调与关键参数中草药叶片数据通常只有几百到几千张从零训练一个 CNN 效果一般很差。最常见的可靠方案是使用在 ImageNet 上预训练的 ResNet18冻结前面几层只微调最后一两个 block 和全连接层。import torch import torch.nn as nn from torchvision import models, transforms from torch.utils.data import DataLoader transform_train transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) transform_val transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) train_set LeafDataset(dataset/train, transform_train) val_set LeafDataset(dataset/val, transform_val) train_loader DataLoader(train_set, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_set, batch_size32, shuffleFalse, num_workers4) model models.resnet18(pretrainedTrue) num_classes len(train_set.class_to_idx) model.fc nn.Linear(model.fc.in_features, num_classes) # 冻结前 7 个 block只微调最后一层和 fc for name, param in model.named_parameters(): if name.startswith(layer4) or name.startswith(fc): param.requires_grad True else: param.requires_grad False optimizer torch.optim.SGD(filter(lambda p: p.requires_grad, model.parameters()), lr0.001, momentum0.9, weight_decay1e-4) criterion nn.CrossEntropyLoss() scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, factor0.5, patience5)参数说明RandomResizedCrop会让模型看到叶片不同尺度和位置的局部比固定中心裁剪对尺寸不一的照片鲁棒得多。RandomRotation(15)对叶片这种没有重力方向依赖的物体很友好但角度别太大超过 30 度会把叶片语义破坏掉。weight_decay1e-4是对抗过拟合的最小配置数据量小的时候可以试着加大到1e-3。学习率我用0.001做微调如果是从零训练这个值通常要降到0.0001判别标准很简单——前 3 个 epoch 训练损失完全不降就把 lr 除以 10 重来。4.3 训练循环与可视化让训练过程透明可复盘训练循环本身不复杂但要把 val 准确率和学习率记录到日志里方便后期绘制曲线。best_acc 0.0 log_lines [] for epoch in range(30): model.train() running_loss 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() optimizer.step() running_loss loss.item() # 验证阶段 model.eval() correct 0 total 0 val_loss 0.0 with torch.no_grad(): for images, labels, _ in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() acc correct / total avg_train_loss running_loss / len(train_loader) avg_val_loss val_loss / len(val_loader) log_lines.append(fepoch:{epoch:02d} train_loss:{avg_train_loss:.4f} val_loss:{avg_val_loss:.4f} val_acc:{acc:.4f} lr:{scheduler.optimizer.param_groups[0][lr]:.1e}) print(log_lines[-1]) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_leaf_model.pth) print(save best model) scheduler.step(avg_val_loss) with open(training_log.txt, w) as f: f.write(\n.join(log_lines))best_acc在 val 上最高时保存模型而不是最后一轮保存这是防止后期过拟合的标准做法。ReduceLROnPlateau的patience5意味着连续 5 轮 val 损失不降就把学习率减半训练到后面你会看到 lr 从 1e-3 一路降到 1e-5这是正常现象。日志写进training_log.txt之后画损失曲线、对比不同超参数结果都靠它别偷懒只打印到终端。5. 中草药叶片识别的四个典型坑现象、原因、对策5.1 验证集准确率很高测试集一塌糊涂数据泄露在作怪现象train 和 val 的准确率都到了 95%但一到 test完全没参与训练的图片上只剩 60%。原因多半是划分数据时没有确保同一株植物的照片只落在同一个集合里。比如你采集时对同一片叶子拍了正反面、不同角度脚本随机划分时这些同源照片可能一张进 train、一张进 val模型相当于已经提前见过答案。解决按个体划分。如果是盆栽植物以植株编号为分组单位保证同一株的所有照片进同一集合如果无法实现至少用imagehash计算感知哈希把重复的近重复图片去重后再划分。5.2 loss 迟迟不下降学习率和标签两头查现象训练了 10 个 epochtrain loss 在 2.0 左右纹丝不动val 准确率还不到 20%。原因最常见的是学习率偏大导致损失函数在陡峭区域震荡另一个可能是标签错位——类别字典生成时排序混乱图片和标签对不上模型学到的东西全都是错的。解决先把学习率直接降到 1e-4 试跑 3 个 epoch如果还是不动打印train_set[0]对应的图片和标签人工看一眼确认图像路径和类别名是否匹配。我遇到过一次典型的翻车文件名是银杏_002.jpg但脚本按str(path)排序时把银杏_010.jpg排到了银杏_002.jpg前面映射编号就乱了肉眼扫一眼样本网格图才发现。5.3 类别字典有时候对不上推理阶段完全混乱现象训练时准确率不错部署到新机器上用torch.load加载模型后预测一个银杏叶图片输出概率最高的类别名是薄荷叶。原因训练时用的类别顺序是sorted(os.listdir(...))推理时换了机器或换了代码在线扫描目录如果目录里混入了缓存文件或者系统大小写排序规则不同顺序就会变。解决训练时把class_indices.json作为唯一可信源推理阶段直接读取不要再去扫描目录重新生成。我之前就吃过这个亏训练时手动把某个异常文件夹重命名为_backup排序整体后移一位类别的索引全换了模型分数全对但语义全错。所以类别字典文件生成后永远别手动改目录结构要加数据也是往已有类别里加。5.4 图像尺寸不统一GPU 显存利用率忽高忽低现象同一批数据有时 batch 跑到一半就 CUDA OOM有时又没事。原因Resize(256) CenterCrop(224)如果放在 DataLoader 里大部分设备没事但如果你在__getitem__里临时做复杂的 transform每张图的尺寸在进入模型前各不相同PyTorch 的 dataloader 会自动 padding 到同 batch 最大尺寸那个最大尺寸基本都偏大显存就爆了。解决统一在 transform 里给定Resize(224, 224)或先用脚本把所有图预先处理成 224×224 存好不要在训练时动态多做一步耗时又占显存的操作。类叶片这种边缘纹理重要的图直接Resize(256)再CenterCrop(224)是最稳的别用RandomResizedCrop生成太小的 crop 分辨率会让叶脉细节全部丢失。5.5 高准确率是错觉模型靠背景认出叶片现象换了一批背景完全不同的测试照片准确率暴跌到 30%。原因采集时薄荷叶多用绿色背景拍银杏叶惯用白色背景模型学到的是背景色域而不是叶片本身。中草药叶片照片大多来自户外、桌面、纸质标本不同环境背景差异极大这个问题尤其严重。解决训练时加入背景增强——随机改变亮度和对比度甚至用RandomResizedCrop偶尔只裁剪出叶片局部更进一步可以做简单的前景分割把叶片从背景中抠出来统一贴到灰底上再训练。我在实际项目中加了一个RandomApply([transforms.ColorJitter(brightness0.3, contrast0.3)], p0.5)val 准确率没变但野外新拍的测试准确率提高了约 20 个百分点这个增强在叶片识别里强烈建议加上。6. 用混淆矩阵和误判样本把准确率钉在业务语义上准确率只是一个数字真正要交付给中医药场景使用时你得知道模型到底在哪些类别之间混淆。比如艾草和益母草叶片形态接近混淆矩阵能一针见血地指出这个高错误率组合然后决定要不要为这两个类单独补拍数据。验证阶段保存混淆矩阵import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix all_preds [] all_labels [] all_paths [] model.eval() with torch.no_grad(): for images, labels, paths in test_loader: images images.to(device) outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().tolist()) all_labels.extend(labels.tolist()) all_paths.extend(paths) cm confusion_matrix(all_labels, all_preds) plt.figure(figsize(12, 10)) plt.imshow(cm, interpolationnearest, cmapBlues) plt.colorbar() classes list(train_set.class_to_idx.keys()) plt.xticks(range(num_classes), classes, rotation90) plt.yticks(range(num_classes), classes) plt.savefig(confusion_matrix.png, dpi150)拿到混淆矩阵后我习惯把错误样本直接打印出来# 找出预测错的样本按真实类预测类分组打印前 5 张 from collections import defaultdict errors defaultdict(list) for label, pred, path in zip(all_labels, all_preds, all_paths): if label ! pred: errors[(classes[label], classes[pred])].append(path) for (true_cls, pred_cls), paths in errors.items(): print(f{true_cls} - {pred_cls}: {len(paths)} 张) for p in paths[:5]: print( , p)打印后你会很清楚地看到错误往往集中在某个具体拍摄条件上——比如所有错图都是侧面拍的叶片。这时我会针对性地调整数据增强增加侧视角的旋转范围或者干脆删掉这些极度相似的干扰样本。这也是我这些年养成的习惯每次训练完先看混淆矩阵再看误判路径列表最后再决定下一步是改模型还是补数据。很多时候改数据比换模型有效得多。这个流程走下去中草药叶片识别的项目才能从跑通走到可信。希望这套从数据划分到误判追溯的完整链路能帮你在类似课题上少走几条弯路。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →