用PyTorch训练花卉识别模型:64类数据集与37种CNN实战

用PyTorch训练花卉识别模型:64类数据集与37种CNN实战 简介面向深度学习图像分类入门与实战这份资料整合了花卉识别专用数据集与多模型训练代码。数据集涵盖64种花卉共32000张224×224彩色图像其中25600张用于训练、6400张用于测试图像均为手机实地采集较网络爬虫图更贴近真实场景。配套训练代码基于卷积神经网络实现37种主流分类模型覆盖ResNet、VGG、Inception、MobileNet、DenseNet、EfficientNet、SqueezeNet等系列方便对比不同架构在花卉识别任务上的效果。整个压缩包约194.92MB包含2000个文件以1919张jpg图像为主另有39个txt说明文件、20个Python训练脚本及22个pyc编译文件可直接用于模型训练与推理。已有1645人学习/下载。对于需要一套干净、成体系的花卉图像分类基准数据的开发者这份资源同时提供了数据与代码省去自行采集标注与搭建模型的流程适合快速上手图像分类项目。1. 手机相机图库中的64类花卉数据集比爬虫图更值得训练做花卉识别第一反应往往是去找 ImageNet 里那几个 flower 类别但一旦把模型放到手机相册里拍的实景花朵上背景、光照、拍摄角度全变了精度掉得很快。这份资源把 64 种常见花卉整理成 32000 张 224×224 彩色图训练集 25600 张、测试集 6400 张图片来自手机实地采集而不是网络爬虫文件名形如024-001-03397.jpg前三位就是类别编号。配套的深度学习训练代码覆盖 ResNet、VGG、Inception、MobileNet、DenseNet、EfficientNet、SqueezeNet 七大系列的 37 种图像分类模型对做深度学习和模型训练的人来说既能当细粒度分类练手项目也能快速验证预训练模型在真实拍摄条件下的泛化能力。2. 数据集解包64类32000张图像的结构、标签与训练堵点原始下载包里的 jpg 是平铺存放的目录里看不到类别文件夹。如果直接把所有 jpg 丢给torchvision.datasets.ImageFolder它会认为每个文件单独一个样本但没有类别归属所以第一件事不是急着写模型而是重建目录结构让训练代码能正确读取标签。2.1 文件名三段式类别编号、样本编号和原始编号024-001-03397.jpg这种命名并不是随机字符串。按常见约定第一段024是类别 ID第二段001是样本编号第三段03397是原始图片在手机相册里的编号。训练时真正用到的是第一段后两段只用于追溯原始照片。如果你的数据集会按类划分通常会把024对应的目录名写成024类别中文名由另一个label.txt维护。先用一段脚本把平铺文件移动到按类别分好的目录里from pathlib import Path import shutil src Path(flower_224) # 原始 jpg 平铺目录 dst Path(flower_by_class) # 按类别分好的目标目录 dst.mkdir(exist_okTrue) for img in src.glob(*.jpg): cls_id img.stem.split(-)[0] # 文件名第一段是类别 ID例如 024 target_dir dst / cls_id target_dir.mkdir(parentsTrue, exist_okTrue) shutil.copy(img, target_dir / img.name)这里img.stem返回不含后缀的文件名split(-)[0]取出第一段024。不要把split(-)[1]当作类别那是样本编号用了会直接多出 64 倍的错误目录。parentsTrue解决多级目录不存在的问题exist_okTrue防止重复执行时抛FileExistsError。更稳妥的做法是同时生成label.txt把024映射成“月季”这类可读名称模型训练脚本只读数字 ID展示结果时才查表。2.2 训练集与测试集划分三个检查点必须过一遍资源给出的是训练集 25600 张、测试集 6400 张的 8:2 划分但下载后首先要核对三个数字类别数是不是 64样本总数是不是 32000每个类别是否都存在足够样本。如果包内已经分好train/和val/直接使用如果还是按类平铺建议手动切分。核对类别数和样本数的脚本可以这样写from PIL import Image from pathlib import Path root Path(flower_by_class) for cls_dir in sorted(root.iterdir()): imgs list(cls_dir.glob(*.jpg)) print(f{cls_dir.name}: {len(imgs)} 张) # 抽检尺寸不用解完所有 32000 张 for p in imgs[:5]: w, h Image.open(p).size assert (w, h) (224, 224), f{p.name}: ({w}, {h})这段代码只解每类前 5 张图避免全量解码带来的 IO 开销。手机图片容易携带 EXIF 旋转信息某些图像解码后宽高会和实际旋转后的尺寸对调抽检可以提前发现这类问题。类别数少于 64 或某个类只有几十张时需要重点关注后面 5.4 节会专门说单类样本不足的处理。确认无误后按类别切分训练集和验证集from sklearn.model_selection import train_test_split import shutil root Path(flower_by_class) train_dir Path(flower_dataset/train) val_dir Path(flower_dataset/val) for cls_dir in root.iterdir(): imgs list(cls_dir.glob(*.jpg)) train_imgs, val_imgs train_test_split( imgs, test_size0.2, random_state42) for p in train_imgs: out train_dir / cls_dir.name / p.name out.parent.mkdir(parentsTrue, exist_okTrue) shutil.copy(p, out) for p in val_imgs: out val_dir / cls_dir.name / p.name out.parent.mkdir(parentsTrue, exist_okTrue) shutil.copy(p, out)test_size0.2对应资源声称的 25600/6400 比例random_state42保证可复现。这里不是全数据集按 8:2 一次性切分而是每个类别内部独立切分这样每个类别在验证集里都占 20%不会出现某个类别在验证集中缺失的极端情况。2.3 数据增强手机采图的背景复杂度比爬虫图更难处理手机拍摄的花卉往往有绿叶背景、曝光不均、轻微失焦。训练时如果只用随机裁剪和水平翻转模型容易记住背景颜色而不是花瓣纹理。常见做法是把增强分成几何扰动和颜色扰动两部分from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0), ratio(0.75, 1.333)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])RandomResizedCrop的scale范围控制在 0.6 到 1.0模拟手机取景时花占画面比例不固定的情况ratio接近正方形避免过度拉伸花瓣形状。RandomRotation(15)处理手机拍摄时经常出现的轻微歪斜。ColorJitter的hue0.05不要调大花卉颜色是分类主特征色调偏移太大会让模型学到错误的颜色分布。验证集只用Resize(256)加CenterCrop(224)不做随机增强。3. 一套代码跑37种CNN统一训练器的模型接口设计拿到数据集之后训练代码的价值就体现出来了。这套代码的入口是一个模型名字符串背后统一封装了七大系列 37 个分类模型换模型不需要改数据加载和训练逻辑。我习惯把模型构建独立成一个model_builder.py避免在训练脚本里堆 37 个 if-else。3.1 用 getattr 动态构造模型再替换分类头TorchVision 的 models 接口把所有网络都挂在同一个命名空间下resnet50、vgg16_bn、mobilenet_v3_large都是可调用对象。用getattr(tv_models, name)能像查字典一样取到对应的构造器之后只需要按系列替换最后的分类层。import torch.nn as nn import torchvision.models as tv_models def replace_head(model, name, num_classes): if name.startswith(squeezenet): in_ch model.classifier[1].in_channels model.classifier[1] nn.Conv2d(in_ch, num_classes, kernel_size1) return model if name.startswith(resnet) or name.startswith(inception): model.fc nn.Linear(model.fc.in_features, num_classes) elif name.startswith(densenet): model.classifier nn.Linear(model.classifier.in_features, num_classes) else: # vgg / mobilenet / efficientnet 都在 classifier[-1] layer model.classifier[-1] model.classifier[-1] nn.Linear(layer.in_features, num_classes) return model def build_model(name, num_classes64, pretrainedTrue): weights DEFAULT if pretrained else None model getattr(tv_models, name)(weightsweights) return replace_head(model, name, num_classes)重点说明几点。第一weightsDEFAULT是新版 TorchVision 的写法等价于自动选择该模型在 ImageNet 上的最佳预训练权重比旧版pretrainedTrue更显式。第二SqueezeNet 分类头是Conv2d而不是全连接层所以替换时要用nn.Conv2d(..., kernel_size1)如果硬塞一个nn.Linearforward 阶段会直接报维度错误。第三replace_head里没有覆盖所有细节比如 InceptionV3 的辅助分类器实际使用时通常会把aux_logitsFalse或保留训练时单独计算 aux loss下面单独说。3.2 Inception 系列的特殊处理输入尺寸和输出结构都不同InceptionV3 官方输入是 299×299不是 224×224。直接在 224 图上跑前面的卷积层会因为空间尺寸不足抛出异常。如果要用inception_v3数据增强里的RandomResizedCrop要改成 299验证集也要对应CenterCrop(299)。另一个坑是训练时 InceptionV3 的 forward 返回的是InceptionOutputs对象而不是普通 Tensor。训练循环里必须取.logits否则CrossEntropyLoss会报TypeError。建议在训练脚本里用一段防御性代码from torchvision.models.inception import InceptionOutputs outputs model(images) if isinstance(outputs, InceptionOutputs): outputs outputs.logits验证或推理时把model.aux_logits False可以直接让模型输出纯 logits省去判断。这个细节是实际跑 Inception 时最常见的卡点不是模型本身难调而是接口形态导致。3.3 七大系列模型选型不只看参数量下面是 37 个模型在分类任务里的一般选型倾向按系列整理成表方便决定先用哪个模型做 baseline模型系列代表模型输入尺寸参数量特点这个数据集上的倾向ResNetresnet18 / resnet50224适中最稳定的 baseline先跑它VGGvgg16_bn / vgg19_bn224大特征可视化方便训练慢Inceptioninception_v3299大多尺度感受野需要改输入尺寸MobileNetmobilenet_v3_large224小手机端部署首选DenseNetdensenet121 / densenet201224小到中特征复用强小数据比 ResNet 稳EfficientNetefficientnet_b0 / b1224小精度/算力性价比高SqueezeNetsqueezenet1_0 / 1_1224极小嵌入式原型验证花卉分类考验的是细粒度特征DenseNet、EfficientNet 通常比 MobileNet 更容易收敛。但如果目标是部署到手机端MobileNet 的推理速度优势不可替代。37 种模型自由选择不等价于每轮都换模型建议先用 ResNet50 跑通全流程再在同一系列内换深度最后对比不同系列在同一超参下的表现。3.4 超参基线预训练模型不要用太大的学习率默认超参我一般这样设optimizer torch.optim.SGD( model.parameters(), lr0.001, momentum0.9, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxepochs) criterion torch.nn.CrossEntropyLoss()ResNet、VGG、DenseNet 用 SGD 加余弦退火就足够初始学习率 0.001 对预训练模型是安全值。MobileNet、EfficientNet 这类 BN 层较多的网络SGD 调起来会有抖动换AdamW(lr1e-4)更稳。如果是从头训练而不是加载预训练权重学习率可以提到 0.01但 32000 张图的数据量不足以支撑从头训练出好的细粒度特征不建议这么干。4. 训练与验证命令、指标与断点续训数据集和模型接口准备好之后训练过程本身要能快速复现、及时看指标、失败后能恢复。这里按常见训练入口脚本整理一套完整流程。4.1 训练入口与命令行参数包里的训练代码如果命名为train.py典型启动命令如下python train.py \ --data flower_dataset \ --model resnet50 \ --epochs 40 \ --batch-size 64 \ --lr 0.001 \ --pretrained \ --gpu 0--data指向包含train和val的根目录--model对应模型名--pretrained表示加载 ImageNet 预训练权重--gpu 0指定显卡编号。--batch-size 64在 8GB 显存上跑 ResNet50 已经接近上限如果显存不够优先降 batch size 到 32而不是调小输入分辨率。输入分辨率降低会影响细粒度特征提取属于最后才考虑的手段。4.2 训练循环里最关键的几行代码训练循环本身不复杂但顺序和状态切换经常出错。核心部分可以写成这样import torch from torchvision.models.inception import InceptionOutputs def train_one_epoch(model, loader, optimizer, criterion, device, model_name): model.train() total_loss 0.0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) if inception in model_name and isinstance(outputs, InceptionOutputs): outputs outputs.logits loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) return total_loss / len(loader.dataset)model.train()必须写在循环外否则每批都切换模块状态。optimizer.zero_grad()必须在backward()之前调用否则梯度会跨 batch 累加。loss.item()取出 Python 数值用于统计乘以images.size(0)加权避免最后一个 batch 样本数不足导致平均 loss 偏低。验证函数要放在torch.no_grad()下并且把模型切到 eval 状态torch.no_grad() def validate(model, loader, criterion, device, model_name): model.eval() correct 0 total 0 val_loss 0.0 for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) if inception in model_name and isinstance(outputs, InceptionOutputs): outputs outputs.logits val_loss criterion(outputs, labels).item() * images.size(0) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return val_loss / total, correct / total注意验证时不要写model.train()否则 BN 层会继续更新 running_mean 和 running_var验证指标会虚高尤其是 MobileNet、EfficientNet 这类 BN 密集的模型。每个 epoch 结束后保存最佳模型if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), fcheckpoints/{model_name}_best.pth)如果训练中断建议保存完整训练状态不只是模型权重torch.save({ model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict(), epoch: epoch, }, fcheckpoints/{model_name}_last.pth)恢复训练时依次load_state_dict再继续循环。只保存模型权重的做法在崩溃后要重新调整学习率很浪费。4.3 测试集评估分类报告和混淆矩阵写文件别刷屏训练结束后需要把当前最好的模型加载回来在测试集上输出 top-1 准确率、每类的精确率和召回率以及混淆矩阵。from sklearn.metrics import classification_report, confusion_matrix model.load_state_dict(torch.load(checkpoints/resnet50_best.pth)) model.eval() y_true, y_pred [], [] for images, labels in test_loader: images images.to(device) with torch.no_grad(): outputs model(images) preds outputs.argmax(dim1) y_true.extend(labels.numpy()) y_pred.extend(preds.cpu().numpy()) report classification_report(y_true, y_pred, digits3) with open(resnet50_report.txt, w) as f: f.write(report) cm confusion_matrix(y_true, y_pred)64 类的classification_report在终端里会刷屏写文件更实用。混淆矩阵重点看对角线外的热点比如某两类互相猜错说明这两类在视觉上高度相似后续要考虑合并类或采集更多样本。5. 让花卉模型再涨两点细粒度分类的五个实用技巧5.1 分两段训练先冻结骨干网络只训练新加的分类头用 1e-3 的学习率跑 5 个 epoch然后解冻全部参数学习率降到 1e-4 继续训练。冻结时要把分类头单独拿出来for param in model.parameters(): param.requires_grad False for param in head.parameters(): param.requires_grad True对 ResNet 系列head就是model.fc对 MobileNet 系列则是model.classifier[-1]。这样做能让随机初始化的分类头先稳定下来避免一开始反向传播的梯度太大把预训练特征冲坏。5.2 用 Label Smoothing 代替普通交叉熵64 类花卉里存在大量相似类别硬标签会让模型对训练集的边界过于自信。改一行代码即可criterion torch.nn.CrossEntropyLoss(label_smoothing0.1)label_smoothing0.1表示真实类别的目标概率是 0.9其余 0.1 分摊到其他 63 类。对细粒度分类任务这通常能带来 0.5 到 1 个点的验证集提升。5.3 遇到长尾类别用 WeightedRandomSampler如果检查时发现某些类别明显少于 400 张训练时要避免多数类主导 loss。使用WeightedRandomSampler按类别样本数的倒数加权采样from torch.utils.data import WeightedRandomSampler class_counts [len(list(cls_dir.glob(*.jpg))) for cls_dir in train_dir.iterdir()] weights [1.0 / c for c in class_counts] sample_weights [weights[label] for label in all_labels] sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue)replacementTrue允许重复采样这样每个 epoch 里少样本类也能被反复看到。5.4 手动检查 top-5 命中但 top-1 错误的样本花卉分类的很多错误并没有真的“分错”只是 top-1 压错了顺序。训练结束后遍历测试集把top-1错误但top-5正确对应的图片收集出来看是光照问题、遮挡问题还是两个类真的长得很像。若是后者建议把两个类合并后再训一轮数据集重新标注的成本远低于继续堆模型复杂度。5.5 测试时增强吃掉最后 0.5 个点验证和测试时可以做简单的水平翻转 TTA把多次推理结果平均base_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean, std) ]) flip_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p1), transforms.ToTensor(), transforms.Normalize(mean, std) ]) torch.no_grad() def tta_predict(model, img, device): model.eval() probs [] for t in (base_transform, flip_transform): x t(img).unsqueeze(0).to(device) probs.append(torch.softmax(model(x), dim1)) final torch.mean(torch.cat(probs), dim0) return final.argmax().item()TTA 的意义在于消除测试图片本身的方向偏置。手机拍摄的花卉图没有固定朝向水平翻转后的预测结果如果和原图不一致说明模型对方向的敏感性高于对花型特征的敏感性。注意 TTA 只适合推理阶段不要在验证集上反复调参时使用否则会把 TTA 带来的提升当成模型本身的提升。本文还有配套的精品资源点击获取