ARTICLE DETAIL

资讯详情

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

花卉识别数据集5类实战:从数据划分到迁移学习分类器

花卉识别数据集5类实战:从数据划分到迁移学习分类器 简介这份资源是面向深度学习入门者与计算机视觉爱好者的花卉识别实战资料包围绕5类常见花卉的物体分类任务展开配套TensorFlow代码与作者录制的B站讲解视频帮助读者从零跑通一个完整的图像分类项目。压缩包共2000个文件约596.75MB其中1972张jpg图片构成训练与测试样本14个py脚本负责数据读取、模型搭建与训练流程另有txt说明、xml标注、md笔记和pdf教程覆盖从数据到代码再到文档的完整链路。目前已有14841人学习下载热度较高。读者可借助现成脚本快速复现分类实验结合视频理解卷积网络训练细节并通过标注与说明文件掌握数据组织方式适合课程设计、入门练手或作为迁移学习的基础数据集使用。1. 花卉识别数据集5类从一个压缩包到能跑通的分类器拿到「花卉识别数据集5类-提供代码和教程.zip」这个标题的人十有八九是三种处境之一课程作业要交一个能演示的分类 demo、想入门计算机视觉但不想从爬虫开始、或者手头有个小场景比如自家花圃、花店库存想验证一下迁移学习到底靠不靠谱。这个数据集的核心价值不在于「5 类」这个数字而在于它把数据、代码、教程打包在一起省掉了最耗时的数据采集和清洗环节让你能直接跳到模型训练和调参这一步。但打包好的东西往往有个通病教程写得像流水账代码能跑但不知道为什么这么写换一批图片就翻车。所以这篇不打算复述压缩包里的说明文档而是按一线做分类任务的顺序把「5 类花卉识别」这件事从数据检查、环境搭建、训练脚本、参数调整到排错完整走一遍。适合能看懂 Python 基础语法、装过 PyCharm 或 Anaconda、但没独立做过图像分类项目的人如果你已经训过 ResNet可以直接跳到第 4 章的参数表和避坑部分。2. 先搞清楚这 5 类数据长什么样目录结构、样本量与划分陷阱2.1 解压后先别急着写代码用三行命令摸清数据底细很多人拿到数据集压缩包第一反应是直接torchvision.datasets.ImageFolder一把梭结果训练到一半发现某类只有 30 张图或者验证集里混进了训练集的图。花卉识别数据集5类这种小规模数据类别不平衡和划分泄漏是两个最隐蔽的坑。先解压然后进目录看一眼# 解压到当前目录下的 flowers5 文件夹 unzip 花卉识别数据集5类-提供代码和教程.zip -d flowers5 cd flowers5 # 查看目录树确认是不是按类别分文件夹 find . -maxdepth 2 -type d | sort # 统计每个类别下的图片数量假设图片是 jpg/png for d in */; do echo -n $d: ; find $d -type f \( -name *.jpg -o -name *.png -o -name *.jpeg \) | wc -l; done这三条命令分别解决三个问题find -maxdepth 2 -type d确认目录层级是不是类别名/图片这种 ImageFolder 能直接读的结构第二条循环统计每类样本数如果发现最多的类有 800 张、最少的只有 120 张后面训练时就得考虑加权采样或者数据增强倾斜第三条其实隐含了一个检查——如果图片散落在多层子目录里ImageFolder 会读不到需要先拍平。常见做法是样本数低于 200 的类训练时用WeightedRandomSampler给它更高采样概率样本数差距在 2 倍以内的直接训也不会太差。我一般会先把统计结果记下来后面调参时对照着看。2.2 训练集/验证集/测试集的划分别用随机 split 毁掉评估可信度压缩包里如果已经分好了 train/val那最好直接用。但很多打包数据集只给一个总文件夹需要自己划分。这里有个血泪经验不要用random.shuffle对整个列表切分因为同一株花可能被连拍多张随机切分会让同一株的图片同时出现在训练和验证集里验证准确率虚高到 99%上线就翻车。正确做法是按「来源」分组再划分。如果图片文件名里带有拍摄批次或原始编号按编号分组如果没有至少用sklearn.model_selection.GroupShuffleSplit或者手动按文件名前缀分组。下面是一个可复现的划分脚本import os import shutil import random from pathlib import Path random.seed(42) # 固定随机种子保证可复现 src_root Path(flowers5) dst_root Path(flowers5_split) classes [d.name for d in src_root.iterdir() if d.is_dir()] for cls in classes: imgs list((src_root / cls).glob(*.*)) # 按文件名前缀分组模拟同一株/同一批次 groups {} for img in imgs: key img.stem.split(_)[0] # 假设文件名格式为 批次号_序号.jpg groups.setdefault(key, []).append(img) group_keys list(groups.keys()) random.shuffle(group_keys) n len(group_keys) train_keys group_keys[:int(n * 0.7)] val_keys group_keys[int(n * 0.7):int(n * 0.85)] test_keys group_keys[int(n * 0.85):] for split_name, keys in [(train, train_keys), (val, val_keys), (test, test_keys)]: out_dir dst_root / split_name / cls out_dir.mkdir(parentsTrue, exist_okTrue) for k in keys: for img in groups[k]: shutil.copy(img, out_dir / img.name) print(划分完成按 7:1.5:1.5 分组切分)这段代码的关键在groups字典它把文件名前缀相同的图片归为一组切分时整组进同一个 split避免同源图片泄漏。random.seed(42)保证每次运行结果一致方便复现。比例 7:1.5:1.5 是小数据集的常用配置验证集用来调参测试集只在最后跑一次。如果压缩包已经分好跳过这步但一定要用find确认 val 里的图片没有和 train 重名。3. 环境搭建与训练脚本从零把 5 类花卉分类跑起来3.1 用 conda 建一个干净环境避开版本冲突这个玄学问题图像分类的依赖链比较长PyTorch、torchvision、Pillow、numpy、matplotlib版本不匹配时经常报一些看不懂的错比如RuntimeError: Couldnt load custom C ops。最省事的办法是单独建环境别在 base 里折腾。下面这套命令在 Windows 和 Linux 下都验证过conda create -n flowers5 python3.9 -y conda activate flowers5 # 安装 PyTorchCPU 版本够用有 NVIDIA 显卡的换成对应 CUDA 版本 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install numpy pillow matplotlib scikit-learn tqdm # 验证安装 python -c import torch, torchvision; print(torch.__version__, torchvision.__version__)参数说明python3.9是兼容性最好的版本3.10 以上某些旧版 torchvision 会出问题--index-url指定 CPU 版下载源如果你有显卡去 PyTorch 官网查对应 CUDA 版本的安装命令替换即可。装完一定要跑最后那行验证能打印出版本号才算成功。如果报找不到 msvcp140.dll装一下 Visual C Redistributable 就行这是 Windows 上的老毛病。3.2 训练脚本迁移学习 数据增强的最小可跑版本5 类花卉、每类几百张图从零训 CNN 基本没戏迁移学习是标准解法。用 torchvision 自带的 ResNet18 预训练权重替换最后一层全连接为 5 输出只训分类头 微调后面几层。下面是完整可跑的脚本import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms, models from tqdm import tqdm # 数据增强训练集用随机裁剪翻转颜色抖动验证集只做 resize 和归一化 train_tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) data_dir flowers5_split train_ds datasets.ImageFolder(f{data_dir}/train, transformtrain_tf) val_ds datasets.ImageFolder(f{data_dir}/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers2) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers2) device torch.device(cuda if torch.cuda.is_available() else cpu) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 5) # 替换为 5 类输出 model model.to(device) criterion nn.CrossEntropyLoss() # 分类头用大学习率主干用小学习率 optimizer optim.Adam([ {params: model.fc.parameters(), lr: 1e-3}, {params: [p for n, p in model.named_parameters() if not n.startswith(fc)], lr: 1e-4} ]) scheduler optim.lr_scheduler.StepLR(optimizer, step_size7, gamma0.1) best_acc 0.0 for epoch in range(15): model.train() for imgs, labels in tqdm(train_loader, descfEpoch {epoch1}): imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() loss criterion(model(imgs), labels) loss.backward() optimizer.step() scheduler.step() # 验证 model.eval() correct total 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) pred model(imgs).argmax(1) correct (pred labels).sum().item() total labels.size(0) acc correct / total print(fEpoch {epoch1} val_acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_flowers5.pth) print(f最佳验证准确率: {best_acc:.4f})逻辑说明RandomResizedCrop(224, scale(0.7,1.0))让模型看到不同尺度的花ColorJitter模拟不同光照这两项对小数据集泛化帮助最大。优化器用了参数分组分类头 lr1e-3、主干 lr1e-4因为预训练权重已经很好大改会破坏特征。StepLR每 7 个 epoch 降一次学习率防止后期震荡。best_flowers5.pth只保存验证集最好的那次避免过拟合模型被留下。参数怎么改如果显存不够把batch_size降到 16 或 8如果验证准确率一直上不去先把RandomResizedCrop的 scale 下限从 0.7 提到 0.8减少过度裁剪如果训练集准确率远高于验证集说明过拟合加weight_decay1e-4到优化器里。3.3 推理与单张图片测试确认模型真的能用训练完不能只看准确率数字得拿一张没见过的图跑一遍。下面这段代码加载保存的权重对单张图片输出预测类别和置信度from PIL import Image import torch.nn.functional as F model.load_state_dict(torch.load(best_flowers5.pth, map_locationdevice)) model.eval() class_names train_ds.classes # [daisy, dandelion, rose, sunflower, tulip] 之类 img Image.open(test_flower.jpg).convert(RGB) tensor val_tf(img).unsqueeze(0).to(device) with torch.no_grad(): probs F.softmax(model(tensor), dim1)[0] top_prob, top_idx probs.max(0) print(f预测: {class_names[top_idx]} 置信度: {top_prob:.4f}) # 打印所有类别概率方便判断是否模棱两可 for name, p in zip(class_names, probs.tolist()): print(f {name}: {p:.4f})这里用val_tf而不是train_tf因为推理时不做随机增强。打印全部类别概率是个好习惯如果 top1 是 0.45、top2 是 0.42说明模型没把握这张图可能模糊或者属于训练集没覆盖的品种实际部署时要加一个置信度阈值低于 0.6 就返回「无法识别」。4. 参数调整与避坑5 类花卉识别最容易翻车的 5 个地方4.1 避坑一验证准确率 99% 但测试集只有 60%现象训练日志里 val_acc 很快冲到 0.98 以上但拿测试集一跑准确率掉到 0.6 左右。原因最常见的是数据泄漏——同一株花的多张照片被分到了训练集和验证集。花卉数据集里这种情况极多因为采集时往往对一株花连拍。另一个原因是验证集太小几百张图里只有几十张验证波动大。解决按第 2.2 节的分组方式重新划分确保同源图片整组进同一个 split验证集至少占总样本的 15%且每类不少于 30 张。如果重划后 val 和 test 差距缩小到 5 个百分点以内说明之前就是泄漏。4.2 避坑二训练 loss 不下降一直卡在 1.6 左右现象CrossEntropyLoss 从 1.61 开始几个 epoch 后还是 1.5 上下准确率跟随机猜差不多。原因学习率设太大导致梯度爆炸或者数据归一化参数用错。很多人直接ToTensor()就送进模型没有做 ImageNet 的 mean/std 归一化预训练权重会「水土不服」。另一个可能是标签没对上ImageFolder 按文件夹名排序生成 0-4 的标签如果文件夹名有中文或特殊字符顺序可能和你想的不一样。解决先确认transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])这行在训练和验证里都加了然后把学习率降到 1e-4 试 3 个 epoch如果 loss 开始降说明之前 lr 太大最后打印train_ds.classes和train_ds.class_to_idx确认标签映射符合预期。4.3 避坑三某类花总是被识别成另一类现象混淆矩阵显示比如「玫瑰」有 40% 被预测成「郁金香」其他类都正常。原因这两类在颜色或形状上本来就接近而训练集里它们的样本数差距大模型偏向多数类。也可能是数据增强过度把玫瑰的关键特征裁掉了。解决先看这两类的样本数如果差距超过 3 倍用WeightedRandomSampler给少数类更高权重然后把ColorJitter的强度调低brightness/contrast 从 0.2 降到 0.1因为颜色是区分花卉的重要特征抖动太强会抹掉差异最后可以单独对这两类做一轮微调只训分类头学习率 1e-4。4.4 避坑四DataLoader 报 num_workers 相关错误现象Windows 下运行训练脚本报RuntimeError: DataLoader worker (pid xxx) is killed by signal或者直接卡死。原因Windows 的 multiprocessing 和 Linux 不同num_workers 0时如果没把主逻辑放在if __name__ __main__:里会无限递归创建子进程。解决把训练循环包进if __name__ __main__:或者简单粗暴地把num_workers设为 0。设 0 的代价是数据加载变慢但对小数据集影响不大。如果一定要多进程在DataLoader里加persistent_workersTrue也有帮助。4.5 避坑五模型保存了但加载时报 key 不匹配现象model.load_state_dict(torch.load(best_flowers5.pth))报Missing key(s) in state_dict: fc.weight, fc.bias或者Unexpected key(s)。原因保存时用了torch.save(model, ...)保存整个模型加载时又用load_state_dict或者保存的是model.state_dict()但加载前没有先实例化同样结构的模型。解决统一用torch.save(model.state_dict(), path)保存加载时先model models.resnet18(); model.fc nn.Linear(512, 5)建好结构再load_state_dict。如果报Unexpected key(s)里有module.前缀说明保存时用了DataParallel加载时加一行state_dict {k.replace(module., ): v for k, v in state_dict.items()}去掉前缀。5. 把准确率再往上推一档两个我常用的进阶技巧5.1 用测试时增强TTA榨出最后 2-3 个百分点模型训完之后如果测试集准确率卡在 88% 左右上不去可以试试 TTA。思路很简单对同一张测试图做多次不同的变换原图、水平翻转、不同裁剪分别推理后把概率平均。这样相当于让模型「多看几眼」再投票对边界样本特别有效。代码不长def tta_predict(model, img_pil, tf_list, device): model.eval() probs torch.zeros(5).to(device) with torch.no_grad(): for tf in tf_list: tensor tf(img_pil).unsqueeze(0).to(device) probs F.softmax(model(tensor), dim1)[0] return (probs / len(tf_list)).argmax().item() # 定义三种变换原图、水平翻转、中心裁剪放大 tta_transforms [ val_tf, transforms.Compose([transforms.Resize(256), transforms.CenterCrop(224), transforms.RandomHorizontalFlip(p1.0), transforms.ToTensor(), transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])]), transforms.Compose([transforms.Resize(256), transforms.CenterCrop(200), transforms.Resize(224), transforms.ToTensor(), transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])]) ]注意RandomHorizontalFlip(p1.0)在这里是确定性的翻转不是随机增强。三种变换覆盖了位置和镜像变化平均后通常能涨 1-3 个点。代价是推理时间变成 3 倍如果对延迟敏感可以只保留原图和翻转两种。5.2 混淆矩阵 置信度阈值知道模型什么时候「不知道」准确率是个笼统指标真正上线时更关心「模型不确定的时候会不会硬猜」。我习惯在测试集上跑一遍输出混淆矩阵和置信度分布from sklearn.metrics import confusion_matrix, classification_report import numpy as np all_preds, all_labels, all_probs [], [], [] model.eval() with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) out F.softmax(model(imgs), dim1) all_probs.extend(out.max(1).values.cpu().numpy()) all_preds.extend(out.argmax(1).cpu().numpy()) all_labels.extend(labels.numpy()) print(confusion_matrix(all_labels, all_preds)) print(classification_report(all_labels, all_preds, target_namesclass_names)) # 看置信度低于 0.6 的样本占比 low_conf np.array(all_probs) 0.6 print(f低置信度样本占比: {low_conf.mean():.2%})如果低置信度样本占比超过 15%说明模型整体没把握要么加数据要么降低分类难度比如先做二分类再细分。如果某一类的召回率明显低对照混淆矩阵看它被错分到哪一类然后针对性补样本。这套流程走下来你对这个 5 类花卉识别方案的能力边界就心里有数了而不是只盯着一个准确率数字。我自己做这类小数据集分类时最大的教训是别在模型结构上折腾太久ResNet18 够用把时间花在数据划分和增强策略上回报更高。另外每次改完参数一定重新跑一遍完整验证别凭感觉觉得「应该会更好」。希望帮到你。本文还有配套的精品资源点击获取
返回列表