☰
花类识别数据集实战:从解压、清洗到迁移学习训练的完整指南
2026/9/28 2:10:34 网站建设 项目流程

简介:面向图像分类入门与实战练习的花卉识别数据集,整套压缩包适用于机器学习和深度学习初学者快速开展五分类训练与验证。数据共包含4242张花朵图像,划分为洋甘菊、郁金香、玫瑰、向日葵、蒲公英五个类别,每类约800张;图片来自数据流、谷歌图片与Yandex图片,分辨率约320×240像素且比例不一,更贴近真实采集场景,可用作迁移学习、数据增强等进阶实验的练手数据。资源包约449.82MB,文件记录共2000项,主体为jpg花朵图片,另含4个Python脚本、2个pyc与1个txt说明文件;脚本可辅助完成数据读取、图片预处理与训练集/测试集划分,适合直接嵌入分类实验流程。目前已有606人学习下载,适合需要花卉图像数据开展分类实验、课程设计或算法练手的开发者。

1. 花类识别数据集.zip:别急着解压,先想清楚你要拿它干什么

花类识别数据集.zip,听起来就是一个压缩包,解压、扔进训练脚本、跑个准确率,完事。但我在实际项目里见过太多人卡在第一步:解压出来的东西跟想象完全不一样,标注格式是花的,目录结构是乱的,图片里还混着水印和表情包。这个数据集的真正价值不在那几百兆图片里,而在它逼你先把数据工程的基本功过一遍——目录怎么组织、标签怎么编码、样本怎么划分、脏数据怎么清洗,这些才是花类识别模型能不能落地的分水岭。适合谁?适合要做图像分类但不想从零爬数据的人,也适合想拿一个干净基准验证自己数据管道的人。不适合谁?不适合指望解压就能训练出生产级模型的人,那是另一套工程。

2. 解压与体检:跑通训练前,先花十分钟看清这个 zip 的真实结构

2.1 解压并检查文件布局,用一棵目录树定位标注格式

拿到花类识别数据集.zip,第一步不是写训练代码,而是先建立一个无菌的检查环境。我一般会在 Linux 服务器或 WSL2 里操作,因为后面所有统计命令都是 bash 风格的。解压命令很简单,但有几个参数值得较真:

mkdir -p /data/flower && cd /data/flower unzip -q ../花类识别数据集.zip -d raw/ ls -la raw/

逻辑说明:-q是安静模式,避免成千上万条解压日志刷屏;-d raw/指定解压目标目录,而不是把 zip 里的内容直接摊在当前目录,这样后续想删掉重来也干净。解压后先看ls -la,确认顶层是单个文件夹还是散落一堆文件。

接下来最关键的是绘制目录树。tree命令不是所有系统都有,可以用find代替:

find raw/ -maxdepth 2 -type d | head -50 find raw/ -maxdepth 2 -type f | head -20

参数说明:-maxdepth 2只往下看两层,避免把几百个类别文件夹全部打爆屏幕;-type d只看目录,-type f只看文件。这两个命令能在 10 秒内告诉你:数据集是「train/val 按类别分子目录」的 ImageFolder 结构,还是「所有图片平铺 + 一个 CSV 标注文件」的平面结构。

常见的数据集结构有两种。第一种是经典的train/类别名/图片.jpg,对应 PyTorch 的torchvision.datasets.ImageFolder,几乎是零成本接入;第二种是images/xxx.jpg + labels.csv,CSV 里写着文件名和类别 ID 的映射。如果你的 zip 里两者都有,比如all/目录加labels.txt,那就要注意了——这通常是别人从某个竞赛平台搬下来重新打包的,标注格式可能带着平台特有的编号体系,后面要重点核对。

2.2 统计类别数与样本量,算出每个类的图片数量分布

结构看清之后,立刻要做的是数量统计。花类识别数据集最怕的不是图片少,而是类别极度不均衡。比如“玫瑰”有 5000 张,“蒲公英”只有 100 张,直接用原始分布训练,模型的预测会严重偏向样本多的类别。

# 统计每个类别目录下的图片数量 find raw/ -mindepth 2 -maxdepth 2 -type f -name "*.jpg" | \ sed 's|/[^/]*$||' | sort | uniq -c | sort -nr | head -30

逻辑说明:-mindepth 2 -maxdepth 2限定只统计二级目录下的文件,假设目录结构是raw/类别/图片.jpg;sed 's|/[^/]*$||'把文件名部分去掉,只保留目录路径;uniq -c按目录聚合计数;sort -nr按数量降序排列。这样一眼就能看出哪些类别是富样本、哪些是贫样本。

如果标注是 CSV,用 Python 统计更稳:

import pandas as pd df = pd.read_csv('raw/labels.csv') print(df['label'].value_counts().head(30))

参数说明:value_counts()返回每个类别的频次。这一步你会发现数据集名称里虽然有“花类识别”,但具体到类别定义可能很宽泛——有的数据集把“花”分成几十个物种,有的只分“玫瑰、向日葵、郁金香”这种粗粒度。这一点直接影响模型容量选择,后面会讲。

2.3 抽样看图和检查元数据,排除损坏文件和标注错位

统计完数量,必须抽图看内容。这是整个流程里最容易被跳过但后果最严重的步骤。用 Python 写一个快速抽样脚本,把每个类别的首尾各抽一张拼成 contact sheet:

import os from PIL import Image import matplotlib.pyplot as plt root = 'raw/' class_dirs = sorted([d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]) fig, axes = plt.subplots(len(class_dirs), 3, figsize=(12, len(class_dirs) * 1.5)) for i, cls in enumerate(class_dirs): imgs = sorted(os.listdir(os.path.join(root, cls)))[:3] for j, img_name in enumerate(imgs): img = Image.open(os.path.join(root, cls, img_name)).convert('RGB') axes[i, j].imshow(img) axes[i, j].axis('off') axes[i, 0].set_ylabel(cls, fontsize=8) plt.tight_layout() plt.savefig('contact_sheet.png', dpi=100)

逻辑说明:sorted保证每个类别的抽样顺序一致;convert('RGB')强制转成 RGB,避免灰度图或带 alpha 通道的 PNG 在后面训练时引起通道数不匹配;axes[i, 0].set_ylabel(cls)在每行左侧标注类别名。跑完这个脚本,打开contact_sheet.png扫一眼,2 分钟能发现大部分问题。

这一眼看过去要确认三件事。第一,图片内容是否和类别名匹配——有没有“玫瑰”目录里混着月季甚至玫瑰花束包装纸的;第二,图片是否来自同一分布——有没有大量网络截图、带水印的商品图、带 UI 界面的手机截图;第三,有没有明显损坏的图片——PIL 打开时报错或解码后全黑的。最后一项可以用一个快速脚本批量验证:

python -c " from PIL import Image import os, sys root='raw/' bad = [] for dirpath, _, files in os.walk(root): for f in files: try: with Image.open(os.path.join(dirpath, f)) as im: im.verify() except Exception as e: bad.append((os.path.join(dirpath, f), str(e))) print('损坏文件数:', len(bad)) for b in bad[:10]: print(b) "

参数说明:Image.verify()只检查文件完整性,不加载像素数据,速度很快;注意 verify 之后不能直接用它做后续处理,要重新Image.open()。这段脚本对几千张图也要不了几秒,任何返回非空的输出都要记录并准备清洗。

3. 从原始图片到可训练样本:目录重排、标签编码与可复现的数据划分

3.1 重排目录为 ImageFolder 标准结构,解决 train/val 混合的问题

绝大多数花类识别数据集在 zip 里不会贴心地替你分好 train/val/test。就算分了,也可能只是按文件名前缀或者干脆全部混在一起。为了让后续训练代码能直接复用torchvision.datasets.ImageFolder,我一般会先写一个重排脚本,把数据集统一成data/train/类别/图片.jpg和data/val/类别/图片.jpg的结构。

#!/bin/bash SRC="raw/" DST="data/" mkdir -p ${DST}/train ${DST}/val # 按 8:2 比例划分每个类别 for cls_dir in ${SRC}*/; do cls=$(basename "$cls_dir") mkdir -p ${DST}/train/$cls ${DST}/val/$cls imgs=("$cls_dir"*.jpg) total=${#imgs[@]} val_count=$((total * 2 / 10)) # 先排序保证可复现性 IFS=$'\n' sorted=($(sort <<<"${imgs[*]}")); unset IFS for ((i=0; i<${#sorted[@]}; i++)); do if (( i < val_count )); then cp "${sorted[$i]}" ${DST}/val/$cls/ else cp "${sorted[$i]}" ${DST}/train/$cls/ fi done done echo "划分完成: $(find ${DST}/train -name '*.jpg' | wc -l) 张训练图, $(find ${DST}/val -name '*.jpg' | wc -l) 张验证图"

参数说明:val_count=$((total * 2 / 10))表示每个类别固定取前 20% 作为验证集,而不是全局随机抽样,这对类别不均衡的数据集至关重要——保证每个类在训练和验证里都出现。排序那步用IFS=$'\n' sorted=($(sort ...))是为了按文件名排序,避免文件系统返回顺序不一致导致每次划分结果不同。

为什么要复制而不是移动原文件?因为原始 zip 可能还要留着做其他实验,复制一份相当于制造“后悔药”。但这个脚本有个隐患:它假设每张图都是.jpg后缀。如果数据集混着.png或.jpeg,这里的*.jpg通配符会把它们漏掉。稳妥做法是用find -type f遍历所有图片扩展名,或者统一用 Pillow 转换。到这里,数据管道的第一个标准化产物就有了——一个能被 ImageFolder 直接加载的目录树。

3.2 用 Python 脚本完成标签编码,生成 id 到类别名的映射文件

目录结构定下来后,类别名本身就是标签,但中文字段名或特殊字符在深度学习框架里容易惹麻烦。训练代码里通常用整数索引做 label,类别名字符串只用于最终的 confusion matrix 展示。所以需要生成一个class_names.txt:

import os train_root = 'data/train' classes = sorted([d for d in os.listdir(train_root) if os.path.isdir(os.path.join(train_root, d))]) with open('class_names.txt', 'w', encoding='utf-8') as f: for idx, cls in enumerate(classes): f.write(f'{idx}\t{cls}\n') print(f'共 {len(classes)} 个类别,映射已保存到 class_names.txt')

逻辑说明:sorted()确保类别索引稳定,不随文件系统遍历顺序变化;idx从 0 开始,与 PyTorch 的 CrossEntropyLoss 默认类别索引一致。这个文件是训练和推理共用的“单一事实来源”——模型输出整数,推理脚本查这个文件转成中文类别名。

这里有个容易忽略的点:如果原始 zip 里的类别名带了空格或括号,比如Rose (red),一定要在映射文件里原样保留,同时确认目录名也一致。我见过有人手改映射文件但忘了改目录名,训练到一半才发现 FileNotFoundError,这个错位非常隐蔽。

3.3 设定固定随机种子,把数据加载器配置成可复现实验

目录和标签都就位,接下来要保证“同一份代码跑两次结果一致”。深度学习训练有大量随机性——数据 shuffle、权重初始化、数据增强的随机裁剪,如果不固定随机种子,前后两次实验就失去了可比性,调参就变成玄学。

import random, numpy as np, torch def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) # 让 cuDNN 使用确定性算法 torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False set_seed(1213)

参数说明:torch.backends.cudnn.deterministic = True会让 cuDNN 选择确定性卷积算法,虽然可能慢一点但结果可复现;benchmark = False关闭自动调优,否则同一个模型在不同批次大小下可能选不同算法。1213这个种子值没有特殊含义,但一旦固定下来就不要轻易改,否则前面所有调参记录都作废。

数据加载器方面,DataLoader的shuffle=True依赖内部随机数生成器,PyTorch 提供generator=torch.Generator().manual_seed(seed)做精细控制。多进程 worker 的随机性更难完全复现,但这已经足够了——训练和验证的划分是确定的,这就保证了准确率的波动主要来自模型本身,而不是数据管道。

4. 训练一个花类识别基线模型:选对网络、设对参数、跑通最小闭环

4.1 为什么迁移学习是花类识别最快见效的路径

花类识别属于细粒度图像分类的入门级变体——类别之间有区分度,但不像“鸟种识别”那样需要数羽毛纹理。处理这种任务,从零训练一个 CNN 是最坏的选择,因为数据集规模通常只有几千到几万张,而花类图片的纹理、颜色、形态特征恰恰是 ImageNet 预训练模型已经学过的底层模式。

常见做法是采用迁移学习:加载 ImageNet 预训练的 ResNet 或 EfficientNet,冻结前几层,只微调最后几层和分类头。我一般会优先试resnet18或resnet34,原因很实际——训练周期短,显存占用小,迭代快,而且花类识别任务通常不需要 ResNet50 那种容量就能达到 95% 以上准确率。如果数据集类内差异很大,比如同一种花有盛开、含苞、枯败等多种状态,再考虑升级到 EfficientNet-B3 或 ConvNeXt-Tiny。

模型的容量选择要匹配数据规模。resnet18有 1100 万参数左右,适合 5000-20000 张图的规模;如果每个类别只有一两百张,用更小的mobilenet_v3_small反而更稳,不容易过拟合。训练花类数据集,过拟合是最大的敌人,数据增强比换大模型更有效。

4.2 用 PyTorch 写一个可运行的花类训练脚本:核心参数逐一说明

下面这个脚本是花类识别训练的最小闭环,自带验证集评估和 checkpoints 保存。我尽量保持通用,你可以直接替换数据路径跑起来:

import torch import torch.nn as nn from torch.utils.data import DataLoader, random_split from torchvision import datasets, transforms, models # ---------- 超参数 ---------- BATCH_SIZE = 32 EPOCHS = 30 LR = 1e-3 # 微调阶段用 1e-3,全网络微调用 1e-4 到 3e-4 MOMENTUM = 0.9 WEIGHT_DECAY = 1e-4 SEED = 1213 # ---------- 数据增强 ---------- train_transforms = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_transforms = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # ---------- 数据集 ---------- train_dataset = datasets.ImageFolder('data/train', transform=train_transforms) val_dataset = datasets.ImageFolder('data/val', transform=val_transforms) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True) # ---------- 模型 ---------- model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) num_ftrs = model.fc.in_features model.fc = nn.Linear(num_ftrs, len(train_dataset.classes)) model = model.to('cuda' if torch.cuda.is_available() else 'cpu') criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=LR, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5)

参数说明:RandomResizedCrop的scale=(0.6, 1.0)限制裁剪比例不低于 60%,避免花的主体被切得只剩一角——花类图片的判别信息集中在花心,而不是整张图的背景;RandomRotation(15)的角度不要超过 15 度,旋转太多会让花瓣结构失真;ColorJitter的饱和度扰动对花类特别有效,因为户外拍摄时同一朵花在不同光照下的饱和度差异极大。SGD 配动量是迁移学习的稳妥组合,Adam 适合从零训练,但微调阶段 SGD 的泛化更好。StepLR每 10 轮降一半学习率,适合 30 轮的训练节奏。

4.3 训练循环里加验证集评估:跨轮保存最优模型

训练脚本不能只跑训练,必须每轮结束都在验证集上评估,并且把最优模型单独存出来。否则你跑完 30 轮,最后保存的那一个 checkpoint 可能在某轮已经过拟合了:

best_acc = 0.0 for epoch in range(EPOCHS): model.train() running_loss = 0.0 for inputs, labels in train_loader: inputs, labels = inputs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() # 验证 model.eval() correct = 0 total = 0 with torch.no_grad(): for inputs, labels in val_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) _, predicted = torch.max(outputs.data, 1) total += labels.size(0) correct += (predicted == labels).sum().item() val_acc = 100.0 * correct / total print(f'Epoch {epoch+1}/{EPOCHS} | Loss: {running_loss/len(train_loader):.4f} | Val Acc: {val_acc:.2f}%') # 保存最优模型 if val_acc > best_acc: best_acc = val_acc torch.save({ 'epoch': epoch + 1, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'val_acc': val_acc, 'class_names': train_dataset.classes }, 'best_flower_model.pth') print(f' 已保存新最优模型,验证准确率 {val_acc:.2f}%')

逻辑说明:model.train()和model.eval()切换训练/验证模式——eval 模式关闭 dropout 和 BatchNorm 的统计更新,如果不切,验证集结果会虚高;torch.no_grad()关闭梯度计算,显存占用大幅下降,而且验证不需要反传;torch.max(outputs.data, 1)取预测类别索引,和前面编码的整数对齐。保存 checkpoints 时把class_names也存进去,比单独维护一个文本文件更保险,推理时只要加载 checkpoint 就能知道类别映射。

这个训练循环在单张 RTX 3090 或 4090 上,resnet18+ 32 batch size 大约每个 epoch 5-10 分钟(取决于图片总数),30 轮大概 3-5 小时。如果你的机器只有 CPU 或者只有 4GB 显存的卡,把 batch size 降到 16,模型换成mobilenet_v3_small,时间可以压缩到 1-2 小时,准确率一般也有 90% 以上。

4.4 类别不平衡时的策略:加权损失函数的引入点

如果前面统计时发现某个类别样本极少,比如少于 50 张,那么直接算 CrossEntropyLoss 会让模型把稀有类全部忽略。此时要给 loss 加上类别权重,权重一般取样本数的倒数或倒数的平方根:

import numpy as np from collections import Counter sample_counts = Counter([label for _, label in train_dataset.samples]) total_samples = len(train_dataset.samples) weights = torch.tensor( [total_samples / sample_counts[i] for i in range(len(train_dataset.classes))], dtype=torch.float32 ).to(device) criterion = nn.CrossEntropyLoss(weight=weights)

参数说明:样本总数 / 该类别样本数是最常用的加权方式——样本多的类别权重小,样本少的类别权重大,让 loss 在稀有类上贡献更大。注意train_dataset.samples是 ImageFolder 内存放的(路径, 标签)列表,可以直接取标签做统计。加了权重之后,训练曲线会有明显波动,验证准确率反而可能下降一点,但混淆矩阵里稀有类的召回率会明显提升。

5. 花类识别避坑手册:解压到训练全流程的 5 个真实踩坑点

5.1 图片文件名包含中文或空格,shutil 复制时突然报错

现象:按类别目录重排时,shutil.copy抛FileNotFoundError,但文件明明存在。排查后发现问题出在文件名里夹杂了全角空格和中文括号,shell 脚本里路径没有正确引用。

原因:zip 文件里的文件名是 UTF-8 编码,中文文件名本身没问题,但如果在 bash 里用未加引号的变量传递路径,空格会被当成分隔符拆词。Python 端的os.listdir能识别,但是某些第三方库统计时对特殊字符处理不当。

解决:所有路径变量一律加双引号:cp "${SRC}/${cls}/${img}" "${DST}/train/${cls}/";Python 脚本里统一用pathlib.Path处理路径,不要手动拼接字符串。清洗阶段最有效的手段是把所有文件名重命名为纯 ASCII,比如类名序号.jpg,丑但绝对安全。

5.2 验证集准确率 95%,但真实场景一测就翻车

现象:模型在清洗过的验证集上表现很好(95%+),换到手机随手拍的花上准确率掉到 60%。反复检查代码也没有 bug。

原因:数据集里的图片分布单一——背景干净、主体居中、光线均匀。真实场景里花是嵌在复杂背景里的,有遮挡、有阴影、有虚化。验证集的分布和部署场景不一致,准确率只是“自我感觉良好”。

解决:训练时增强背景扰动,RandomResizedCrop的scale下限调到 0.3,并加入RandomErasing模拟遮挡;另建一个 50-100 张的“野采”测试集,从网上下载不同拍摄风格的花图,单独做一次评估,这才是真实水平的反映。如果野采测试集准确率低于 80%,说明数据集本身和你的落地场景有 gap,要回头补充数据。

5.3 训练时 loss 直接变成 nan,前两步就崩

现象:loss 在第一个 batch 后就变成nan,验证集准确率一直是 0。检查数据没问题,代码也没有明显错误。

原因:最常见是学习率过大,或者输入图片里有全黑的损坏图导致 BatchNorm 计算方差为 0。resnet18预训练模型在 ImageNet 归一化参数下,如果输入值域不对,比如忘记ToTensor()导致输入还是 0-255 的整数,梯度会瞬间爆炸。

解决:先确认预处理管道里ToTensor()和Normalize按顺序执行;再检查有没有全黑或全白图片,用前面 5.2 的 verify 脚本再跑一遍;最后把学习率降到1e-4试试——如果 loss 恢复正常,说明原来的 LR 对这个 batch size 太大了。血泪经验:nan问题 80% 出在数据端,不是模型端。

5.4 同样代码在两个环境跑,验证集准确率差 3 个百分点

现象:同一个脚本,在 A 服务器上训练完验证集准确率 93%,在 B 服务器上复现只有 90%。代码一样、数据一样、超参一样,检查了随机种子也一样。

原因:PyTorch 版本差异导致ResNet18_Weights.IMAGENET1K_V1的下载权重可能有微小差异;cuDNN 版本不同导致卷积算法的浮点舍入不同;数据加载器的num_workers不同导致 shuffle 顺序不同。

解决:训练模型前把torch.__version__、torchvision.__version__、CUDA 版本记录下来,写进训练日志开头;固定随机种子的同时,把DataLoader的generator也固定;如果要求极高,可以用torch.use_deterministic_algorithms(True)强制所有算子确定性,但这会禁用某些不支持确定性的算子,慎用。真正要复现的是实验结论,不是逐比特一致。

5.5 数据集里的“花”和你要识别的“花”不是一回事

现象:用花类识别数据集.zip 训练后,模型对月季、蔷薇、玫瑰的区分度很差,但对菊花、向日葵、郁金香的识别率很高。看起来是数据集本身的类别定义和你的需求错位。

原因:很多公开花类数据集是基于牛津 102 类花卉(Oxford-102)或类似竞赛数据集构建的,类别划分遵循植物学分类,而你的业务诉求可能是按商品名分类——比如“玫瑰”和“月季”在植物学上是不同物种,但电商场景里它们可能归在同一类。

解决:先看class_names.txt里的类别列表,对照你的实际需求逐项核对。如果默认类别定义与你需求不符,合并相近类别、重命名标签都是可行的——直接在data/train目录下用软链接或复制重命名即可,不需要重写训练代码。最怕的是不看类别就拿去用,训出来才发现标签体系对不上。

6. 让模型真正可用:混淆矩阵定位易混类,ONNX 导出提速部署

训练完拿到一个 90%+ 的模型,并不意味着项目结束了。我一般会再做两件收尾的事:一是用混淆矩阵找出模型分不清的类别,二是把 PyTorch 模型导出成 ONNX,让推理速度提升一个档次。

混淆矩阵的实现非常直接,核心就是收集预测结果:

import seaborn as sns import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, classification_report all_preds = [] all_labels = [] model.eval() with torch.no_grad(): for inputs, labels in val_loader: inputs = inputs.to(device) outputs = model(inputs) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) class_names = train_dataset.classes plt.figure(figsize=(12, 10)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.xlabel('Predicted') plt.ylabel('True') plt.xticks(rotation=45, ha='right') plt.tight_layout() plt.savefig('confusion_matrix.png', dpi=150) print(classification_report(all_labels, all_preds, target_names=class_names))

逻辑说明:classification_report会输出每个类别的 precision、recall、F1-score,比只看准确率信息量大得多。比如“菊花”和“蒲公英”在照片里都是白色放射状花瓣,如果这两类的混淆矩阵数值明显偏高,说明模型抓取的特征还不够分,可以考虑针对性增加这两类的样本,或者裁掉图片里过多的背景区域。

关于 ONNX 导出,PyTorch 自带torch.onnx.export,但有几个参数直接影响部署时的表现:

dummy_input = torch.randn(1, 3, 224, 224).to('cpu') model_simple = model.to('cpu').eval() torch.onnx.export( model_simple, dummy_input, 'flower_model.onnx', input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}}, opset_version=13 )

参数说明:dynamic_axes把 batch 维设为动态,这样导出的 ONNX 模型可以接受任意 batch size 的输入,部署到生产环境时不需要固定为 1。opset_version=13比较保守,兼容性最好。导出后建议用onnxruntime验证输出是否和 PyTorch 一致:

python -c " import onnxruntime as ort import numpy as np sess = ort.InferenceSession('flower_model.onnx') input_name = sess.get_inputs()[0].name ort_out = sess.run(None, {input_name: np.random.randn(1, 3, 224, 224).astype(np.float32)}) print('ONNX 输出维度:', ort_out[0].shape) "

到这一步,花类识别数据集的价值才真正闭环——你手里有一个能跑通训练、能定位错误、能导出部署的完整链路。最后说一个我的习惯:每次训练完,我会把epochs、batch_size、lr、最终验证准确率、数据增强配置记成一个experiment.log存进项目目录。这个文件在三个月后回看时比代码本身更值钱——没有哪个参数组合是全能的,记录了什么能跑,比知道什么理论最优更实在。希望帮你避开我踩过的坑。

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

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

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

立即咨询