简介:花卉识别方向的数据集与训练代码组合包,面向计算机视觉初学者、算法工程师及高校相关课题人员,解决花朵分类任务中训练数据不足与模型选型困难的问题。包内包含16种花卉共32000张224×224彩色图片,涵盖千屈菜、射干、曼陀罗、桔梗、狗尾草、秋英、粉黛乱子草、红花酢浆草、黄金菊等常见与观赏植物,每类约2000张,分布均衡,可直接用于训练、验证与测试。配套训练源码基于TensorFlow编写,集合了23种主流图片分类模型,支持根据场景灵活选择网络结构开展对比实验,也便于二次改造。全部文件共110个,以txt标注/说明、py训练脚本、pyc编译文件为主,另有bat启动脚本与jpg示例图,压缩包整体684.45MB,目录层级清晰。已有1499人学习下载,适合需要现成花卉数据集和多种模型基线、快速搭建分类项目或撰写实验报告的开发者。
1. 花卉识别训练源码:从一张花图到可部署模型的最小闭环
做花卉识别的时候,大部分人第一步不是选模型,而是去整理一堆名字乱七八糟的花卉图片集。标题里的“训练源码”和“花卉数据集”其实是同一个项目里的两座大山:数据决定了模型的上限,源码和参数决定了这个上限能不能够到。我见过太多训练日志漂亮、一到户外拍几张真花就翻车的案例,问题往往不在网络结构,而在图片集的标签噪声和训练脚本里的细节。这篇就围绕“花卉识别-花卉数据集-花卉识别训练源码-花卉图片集(02)”这个标题,把数据整理、迁移学习训练、日志判读、避坑和推理部署串成一个可复现的完整闭环,适合准备用自己图片集做花卉分类的工程师和学生照着重跑一遍。
2. 把花卉图片集整理成训练集:目录划分、去重与标签映射
2.1 先用公开花卉数据集验证,再决定是否自制数据集
市面上的公开花卉数据集不算少,常见的有 Oxford 102 Flower 这类按类别组织好的图片集,也有一批按花期、颜色、病虫害细分标注的数据集。这些公开集的好处是标签已经清洗过,有相对标准的类别划分和拍摄场景,适合拿来当“标准答案”跑通训练源码、验证模型有没有选对。坏处是公开集的类别往往和实际业务对不上:你想识别的是自己园区里的十几种花,公开集却按植物学分类给你一百多个类别,视角也不一致。
我一般这么定位:先用公开花卉数据集把训练脚本和超参跑一遍,确认那一套源码在这个任务上能收敛、不翻车,再切换到自己的花卉图片集。自制的图片集优势是贴合业务场景,劣势也很明显——采集来的图片命名混乱、重复度高、同一种花在不同光照和角度下差异悬殊,标签错一张就是一颗老鼠屎。标题里既然出现了“花卉数据集”,就别轻视这一步:训练前花两小时整理目录,能省下训练后两周的排错时间。
2.2 目录划分脚本:按类别拆出 train/val/test
PyTorch 的torchvision.datasets.ImageFolder要求数据按类别目录/图片文件组织,并且会把每个子目录名当成一个类别。常见做法是在数据集根目录下再分train/val/test三个子目录,每个子目录里再放类别文件夹。下面这段脚本能把一堆散乱的类别图片按比例拆开,并保证每个类别都均匀分到三个集合里。
import os import random import shutil from collections import defaultdict # 配置 src_dir = "raw_flowers" # 原始图片根目录,每个子目录一个类别 dst_dir = "flowers_dataset" # 输出目录 train_ratio = 0.7 # 训练集比例 val_ratio = 0.2 # 验证集比例 seed = 42 random.seed(seed) # 收集每个类别的所有图片 class_images = defaultdict(list) for class_name in os.listdir(src_dir): class_path = os.path.join(src_dir, class_name) if not os.path.isdir(class_path): continue for fname in os.listdir(class_path): if fname.lower().endswith((".jpg", ".jpeg", ".png")): class_images[class_name].append(fname) # 按类别独立划分并复制 for class_name, fnames in class_images.items(): random.shuffle(fnames) n = len(fnames) n_train = int(n * train_ratio) n_val = int(n * val_ratio) splits = { "train": fnames[:n_train], "val": fnames[n_train:n_train + n_val], "test": fnames[n_train + n_val:], } for split_name, split_files in splits.items(): out_class_dir = os.path.join(dst_dir, split_name, class_name) os.makedirs(out_class_dir, exist_ok=True) for fname in split_files: src = os.path.join(src_dir, class_name, fname) dst = os.path.join(out_class_dir, fname) shutil.copy2(src, dst)这段脚本的核心逻辑是“按类别洗牌后再切分”。如果先给所有文件全局洗牌再按比例切,某些类的图片可能会全部落到训练集,验证集里缺了这个类,训练日志上的 acc 就会虚高,模型实际泛化能力完全看不见。另一个细节是用了copy2而不是move,这样即使后续发现某个类需要重新划分,原始图片还在,不吃后悔药。
比例参数上,train 占 0.7、val 占 0.2、test 留 0.1 是常见默认值。类别越多、每类图片越少,train 的比例不妨降到 0.8,保证验证集每类至少有 5 到 10 张。数据量特别少的时候(每类不到 20 张),我会只分 train 和 val 两部分,test 直接复用真实场景拍摄的几张照片,而不从这份小数据集里再抠。
2.3 标签映射与不平衡检查:训练前最后一关
ImageFolder会自动生成class_to_idx映射,规则是按子目录名的字母序从 0 编号。这个“自动”既是方便也是隐患:如果你在训练中途改过文件夹名,或者某个子目录里混进了别类的图,映射就悄悄变了。训练前必须花两分钟检查各类别数量分布,脚本如下:
import os from torchvision import datasets data_root = "flowers_dataset" for split in ["train", "val", "test"]: ds = datasets.ImageFolder(root=os.path.join(data_root, split)) # 类别名与数字索引的映射 print(f"== {split} == class_to_idx: {ds.class_to_idx}") # 按类别统计图片数量 class_count = {} for path, idx in ds.samples: class_name = ds.classes[idx] class_count[class_name] = class_count.get(class_name, 0) + 1 for name, cnt in sorted(class_count.items(), key=lambda x: x[1]): print(f" {name}: {cnt}")这段代码做两件事:一是把class_to_idx打出来,人工核对类别名和索引是否和预期一致;二是统计每个类的图片数。如果某个类明显少于其他类,训练时它的 loss 贡献会被大类别淹没。处理方法有几种:给少样本类别做简单过采样(重复读图)、在WeightedRandomSampler里按类别数量反比采样,或者干脆补拍一批该类的图片。数据不平衡在花卉识别里经常被忽略,因为表面上看每一类都有几十张图,但实际拍摄条件下某些花就是难拍到,这类“尾巴类”往往才是上线后误判的重灾区。
3. 选 backbone 和训练框架:迁移学习是花卉识别最快的落地路径
3.1 为什么从 ResNet18 而不是从零训练开始
从零训练一个 CNN 分类网络需要百万张图片作为支撑,花卉图片集的规模通常在每类几十到几百张,从零训练几乎必然过拟合。ImageNet 上预训练过的模型已经学会边缘、纹理、形状这些通用特征,花卉识别需要做的只是在这些特征之上“微调”出一个新的分类头。这个过程在工程上叫迁移学习,也是当前从业者做小数据集图像分类最稳妥的默认方案:代码简单、收敛快、泛化也远比从零训练好。
backbone 选 ResNet18 是基于成本和收益的权衡。ResNet18 只有约 1100 万参数,单张消费级显卡能轻松跑起来,推理速度也快,在 10 到 20 类花卉的常见任务里精度已经够用。只有当你发现模型在验证集上仍然欠拟合、且所有调参手段用尽时,才考虑升级到 ResNet50 或 EfficientNet-B0。输入分辨率也可以从 224 提到 384,这比盲目加大网络层数更直接有效,但显存占用会翻倍。如果目标是识别月季和玫瑰这种细粒度品种,224 分辨率往往不够捕捉花瓣细节,那从一开始就应该用 ResNet50 + 高分辨率输入。
3.2 最小可运行的 PyTorch 训练脚本
下面是一个可复现的花卉识别训练脚本,从读取数据集到保存最优权重一步到位。配合上一章整理好的目录结构,换数据集时只需要改data_root和num_epochs。
import os import json import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, models, transforms # 超参数 data_root = "flowers_dataset" batch_size = 32 num_epochs = 30 lr = 1e-3 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 数据增强与归一化 train_transform = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_transform = 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_ds = datasets.ImageFolder(os.path.join(data_root, "train"), transform=train_transform) val_ds = datasets.ImageFolder(os.path.join(data_root, "val"), transform=val_transform) num_classes = len(train_ds.classes) print("类别数:", num_classes, "类别:", train_ds.classes) train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=4) val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=4) # 加载预训练模型并替换分类头 model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) model.fc = nn.Linear(model.fc.in_features, num_classes) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) # 训练循环 best_val_acc = 0.0 for epoch in range(num_epochs): 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() * images.size(0) epoch_loss = running_loss / len(train_ds) # 每个 epoch 结束做一次验证 model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) val_acc = correct / total print(f"epoch {epoch+1}/{num_epochs} loss={epoch_loss:.4f} val_acc={val_acc:.4f}") # 只保存验证集上表现最好的权重 if val_acc > best_val_acc: best_val_acc = val_acc torch.save(model.state_dict(), "best_flower_model.pth") # 保存类别名,供推理时使用 with open("class_names.json", "w") as f: json.dump(train_ds.classes, f, indent=2)几个关键逻辑点需要展开说明。第一,train_transform和val_transform必须分开:训练集用随机裁剪和颜色扰动,验证集只用固定缩放和中心裁剪,否则验证集的评估会被随机增强干扰,每次跑的 acc 都不一样。第二,model.fc被替换成新的全连接层,这一层的参数从零开始随机初始化,需要更大的学习率去适应;而预训练部分的 backbone 已经收敛,学习率过大会把学到的通用特征摧毁。第三,保存权重时只认验证集上的最优表现,不在测试集上反复挑模型,否则测试集就失去了真实评估的意义。
num_workers=4在 Windows 上偶尔会遇到多进程数据加载报错,如果遇到就直接改成num_workers=0。损失函数CrossEntropyLoss自带 softmax,所以模型的输出不需要再做一次 softmax,直接用 logits 算 loss 更数值稳定。
3.3 训练参数怎么设:lr、batch_size、epoch 与设备的匹配
训练参数不能脱离硬件和数据量单独谈。下面是不同设备上我用过的参考配置,表格里的组合至少能保证“训练跑得起来”:
| 训练场景 | backbone | batch_size | 学习率 | epoch |
|---|---|---|---|---|
| CPU(仅功能验证) | ResNet18 | 8 | 1e-4 | 10-15 |
| 6-8GB 显存显卡 | ResNet18 | 32 | 1e-3 | 30-50 |
| 12GB 以上显卡 | ResNet50 | 64 | 2e-3 | 50 |
| 细粒度品种识别 | ResNet50 + 384输入 | 16 | 1e-3 | 40-60 |
batch_size 直接影响学习率的合理取值。32 是 ResNet18 在常见消费级显卡上的甜点值,显存不够就降到 16 或 8,但不要只降 batch_size 不降学习率,否则梯度噪声偏大、训练不稳定。AdamW 配合weight_decay=1e-4是我在分类任务上的默认组合,比纯 Adam 更好控制过拟合。epoch 的多少要看验证集曲线:花卉小数据集上通常 30 个 epoch 内就能看到明显收敛,超过 50 个 epoch 后边际收益很低,更多是在做无用功。
学习率的设定还有一个误区:很多人微调时直接沿用 ImageNet 训练的初始学习率,例如从头训练用的 0.1,结果第一个 epoch loss 直接飞掉。迁移学习的正确姿势是从1e-3附近开始,如果发现 loss 震荡就降到1e-4。你可以简单地把学习率理解成“模型信任新数据的程度”,数据越脏、类别越细,学习率就要越保守。
4. 跑通训练之后:数据增强、warmup 与训练日志判读
4.1 数据增强:花卉图片最有效的几个操作
花卉识别的图片天然存在三个问题:拍摄尺度差别大、光照条件各不相同、背景复杂多变。数据增强就是为了让模型不依赖这些与花无关的表面特征。以我的经验,最有效的三个操作是RandomResizedCrop、RandomHorizontalFlip和ColorJitter。
RandomResizedCrop随机截取图片的一部分并缩放到目标尺寸,模拟远近不同的拍摄距离。对花卉来说这特别重要,因为同一个品种的特写和远景照片差别非常大。RandomHorizontalFlip几乎不改变花的语义,但能成倍扩充训练样本。ColorJitter模拟一天内不同时段的光照变化,尤其在室外拍摄的图片集里非常有效。
要注意别做过头。见过有人给花卉图片加RandomRotation(45),结果训练时频繁出现倒挂的花,模型被迫学习“转正”而不是学习“这是什么花”。还见过把病害叶片的图片和健康花朵混在一个类别里训练,数据增强反而会把病斑特征放大,让模型学到捷径——这个问题在涉及花卉病虫害数据集时尤其容易踩,后面避坑章还会展开。
4.2 学习率 warmup 与余弦退火:让训练过程更稳
训练初期模型权重刚从预训练迁移过来,直接上大学习率会让 loss 飙升。warmup 的思路是先用小学习率跑几个 epoch,等 loss 稳定下来再把学习率升到设定值。余弦退火让学习率随训练进度从高到低平滑衰减,帮助模型在后期做精细收敛。两者搭配是当前图像分类训练里的常见做法,代码量不大,收益却明显。
from torch.optim.lr_scheduler import SequentialLR, LinearLR, CosineAnnealingLR warmup_epochs = 3 total_epochs = 30 scheduler = SequentialLR( optimizer, schedulers=[ LinearLR(optimizer, start_factor=0.1, total_iters=warmup_epochs), CosineAnnealingLR(optimizer, T_max=total_epochs - warmup_epochs) ], milestones=[warmup_epochs] )这段代码配合上一章的训练循环即可。SequentialLR把两个调度器串起来:前 3 个 epoch 用LinearLR从 10% 的学习率线性升到全量,后面 27 个 epoch 走余弦退火。milestones是切换调度器的节点,必须和warmup_epochs保持一致。注意T_max填的是余弦退火覆盖的 epoch 数,而不是总 epoch 数,写错了学习率曲线会不完整。
4.3 训练日志怎么读:loss、top-1 acc 和过拟合信号
训练脚本每个 epoch 打印一条日志,新手往往只盯 val_acc,忽略了 loss 和 acc 之间的关系。下面是一个典型的小数据集训练过程:
epoch 1/30 loss=1.8563 val_acc=0.6214 epoch 5/30 loss=0.5248 val_acc=0.8027 epoch 10/30 loss=0.2179 val_acc=0.8936 epoch 15/30 loss=0.1205 val_acc=0.9062 epoch 25/30 loss=0.0680 val_acc=0.8849 epoch 30/30 loss=0.0431 val_acc=0.8613前 15 个 epoch 是正常的:训练 loss 从 1.8 降到 0.12,val_acc 同步从 62% 涨到 90%,说明模型在学习。但从第 15 个 epoch 开始,训练 loss 还在继续下降,val_acc 却不升反降,这是典型的过拟合信号。对花卉小数据集来说,这个转折点通常出现在 15 到 25 个 epoch 之间,不用等到 30 个 epoch 跑完才能判断。
发现过拟合后,优先做三件事:一是给数据增强加码,比如提高ColorJitter的幅度或加入RandomAffine的轻微平移;二是给分类头加一点 Dropout;三是把weight_decay从1e-4增大到5e-4。所有手段都试过仍然过拟合,才考虑减少 backbone 的参数量或提前停止训练。val_acc 和 train_acc 相差 5 个百分点以内是正常区间,差超过 10 个点就要警惕。训练日志本身就是黑匣子的唯一窗口,多记录、多对比,比反复改网络结构更能定位问题。
5. 花卉识别训练避坑:5 个让我重跑数据的真实问题
5.1 同一张图同时出现在训练集和验证集:val_acc 虚高不真实
现象:训练日志上 val_acc 一路升到 95% 以上,模型上线后拿真实场景的图片一测,准确率掉到 60% 附近,完全不是训练时的水平。
原因:数据集里存在重复图片。最常见的是从网上采集的花卉图片,同一个文件被改名为不同文件名,或同一朵花的照片被不同来源重复收录。如果这些重复图恰好分别落在 train 和 val 里,模型相当于提前看到了“答案”,val_acc 自然虚高。
解决:划分数据集之前必须先做去重。按文件 MD5 建索引,找出内容完全相同的图片:
import hashlib from pathlib import Path def file_md5(path): h = hashlib.md5() with open(path, "rb") as f: for chunk in iter(lambda: f.read(8192), b""): h.update(chunk) return h.hexdigest() md5_to_path = {} for img_path in Path("raw_flowers").rglob("*"): if img_path.suffix.lower() not in {".jpg", ".jpeg", ".png"}: continue md5 = file_md5(img_path) if md5 in md5_to_path: print(f"重复图片: {img_path} 与 {md5_to_path[md5]}") else: md5_to_path[md5] = img_path这段脚本会打印所有内容相同的图片对。处理方式不是简单地删一张,而是要人工看一眼:如果是同一朵花的重复照片,删掉其中一份;如果是同一类别下的相似但不同照片,MD5 不会判重,说明不在这个问题的范围内。MD5 去重只解决完全相同的文件,无法识别“同一素材被重新压缩或加了水印”的变体,这类情况就得靠感知哈希了,但作为第一道防线已经够用。
5.2 类别标签错乱:class_to_idx 与文件夹名对不上
现象:训练了十几个 epoch,训练 loss 不降或者降得极慢,val_acc 一直在类别数分之一附近徘徊,例如 10 分类时 acc 稳定在 20% 上下。
原因:某个类别文件夹里混入了其他类的图片,或者手工改过文件夹名但没重新生成标签文件。ImageFolder 按字母序自动生成索引,你以为“玫瑰”是 0 号类别,实际 0 号可能是“牡丹”。
解决:训练脚本里加上 2.3 小节的类别检查代码,把class_to_idx和每个类的样例路径打印出来。训练前人工抽查 10 到 20 张图片,确认目录名和内容一致。这个检查 30 秒能完成,但能省下几个小时的无效训练。如果你从网上爬图片时把“花名-来源网站”作为目录名,更要检查一遍,来源网站名混进类别名这种事我见过不止一次。
5.3 推理 transform 与训练不一致:单张图片预测精度骤降
现象:训练脚本里 val_acc 90% 以上,自己写了个单张图片推理脚本,随便拿一张验证集里的图去测,模型给出的预测居然错了,而且置信度很低。
原因:推理脚本里可能只用了Resize(224)或干脆没有做归一化,而训练和验证时用的是Resize(256) -> CenterCrop(224) -> Normalize。输入分布不一致,模型在验证集上学的“经验”完全用不上。
解决:推理时必须复用训练时的val_transform,特别是Resize到CenterCrop的顺序和Normalize的均值方差。最好的做法是把val_transform单独抽成一个函数,训练脚本和推理脚本都从同一个地方导入,而不是在推理脚本里抄一遍。这个坑是最隐蔽的,因为报错不会出现,模型默默给你一个错误的高置信度结果。
5.4 细粒度品种混淆:月季、玫瑰和蔷薇分不清
现象:大类识别不错,比如能把菊花和玫瑰分清,但月季、玫瑰、蔷薇这几个近缘品种互相混,常用的增强和调参手段都救不回来。
原因:这三个品种在图像上差异极小,224x224 分辨率下花瓣结构和叶子形状的特征可能只有几个像素的差别。预训练模型的特征是面向 1000 类通用任务的,在细粒度区分上先天不足。
解决:两条路。一条是提高输入分辨率,把输入从 224 提到 384 或 448,让模型“看得更细”,同时配合更大 batch 或更小学习率;另一条是在工程上调整分类体系,先用粗分类模型定大类,再在类内训练一个细粒度模型,而不是让一个模型做所有事。如果数据量本身就少,提高分辨率的意义也有限,因为模型没有足够的样本去学习那些细微特征。这时候先补数据比调模型更现实。
5.5 病虫害图片混入:模型学到的是病斑而不是花
现象:训练 loss 正常下降,val_acc 也升高,但模型对正常健康花卉的误判率很高,反而对带有病虫害特征的图片“格外有把握”。
原因:采集图片时把带有病斑、黄叶的样本也塞进了对应花类的文件夹。模型在训练中发现“有褐色斑点就是玫瑰”这种捷径特征比花瓣结构更容易区分,于是学到的不是玫瑰的特征,而是病斑的特征。这类问题在整理花卉病虫害数据集时尤其突出——病虫害样本往往具有明显的视觉共性,模型很容易抓住这个共性,忽略真正的花朵结构。
解决:回看数据源,把带明显病害特征的样本单独拿出来。如果业务目标本身就是做病虫害预警,那就应该单独设置“病害玫瑰”这样的类别,而不是混进健康玫瑰里;如果目标只是花种识别,这些样本应直接剔除,宁可少一些训练数据,也不要让标签变得不纯净。判断方法很简单:单独挑一批完全健康的花图做验证,看看模型在这批图上的表现是不是远低于训练时的 val_acc,如果是,就去数据里找“捷径特征”的来源。
6. 从训练到可用:推理脚本、top-k 输出与置信度校准
6.1 单张图片推理脚本:top-3 与置信度一起返回
训练好模型后,下一步是写一个能处理单张图片的推理脚本。注意这里不能直接拿训练脚本改,因为训练脚本里的模型是train模式,会计算 dropout 和 batch norm 的滚动统计量,推理必须切到eval模式。
import torch import json from PIL import Image from torchvision import models, transforms device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 读取训练时保存的类别名 with open("class_names.json") as f: class_names = json.load(f) num_classes = len(class_names) # 重新构建模型结构并加载权重 model = models.resnet18() model.fc = torch.nn.Linear(model.fc.in_features, num_classes) model.load_state_dict(torch.load("best_flower_model.pth", map_location=device)) model.to(device).eval() # 与训练时 val_transform 保持一致 transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) def predict(image_path, topk=3): img = Image.open(image_path).convert("RGB") x = transform(img).unsqueeze(0).to(device) with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1).squeeze(0) topk_indices = probs.topk(topk).indices.tolist() return [(class_names[i], probs[i].item()) for i in topk_indices] # 示例 for name, conf in predict("test.jpg", topk=3): print(f"{name}: {conf:.4f}")这个脚本的class_names必须来自训练时保存的 json 文件,不能重新定义,否则索引对不上。返回 top-3 而不是只返回 top-1,是因为花卉识别里近缘品种本来就难分,把三个候选都交给调用方,让业务侧做最终判断,比模型硬给一个答案更可靠。Image.open(...).convert("RGB")这一步很重要:有些手机拍出的照片是 RGBA 模式或带 EXIF 方向信息,不转 RGB 或不做归一化,预测结果会被这些无关因素干扰。
6.2 置信度校准与阈值:模型说 0.97 也不一定可靠
很多刚跑通训练的人看到softmax输出 0.97 就觉得模型很确定,但实际统计下来,这个 0.97 的预测准确率可能只有 80%。这种“过度自信”在花卉数据集上很常见,尤其是在细粒度品种分类里。原因是训练数据的类别分布和真实场景差异很大,模型拟合的是训练集的分布,而不是真实世界的分布。
把置信度校准回真实概率的常用做法是温度缩放。具体操作是:在验证集上收集模型的 logits 和真实标签,搜索一个最优的温度参数 T,让logits / T经 softmax 后的交叉熵损失最小,然后推理时用调整后的概率作为置信度。T 如果小于 1,说明模型过于保守;大于 1,说明模型过度自信。这个原理不复杂,但很多人不知道——拿到模型后第一件事就是看它输出的概率分布,而不是看提不提供。
校准之后还有一个必须做的动作:设定拒识阈值。常见做法是把 top-1 置信度低于某个值的图片直接判为“无法识别”,而不是硬分到某一个类。阈值的取值可以在验证集上扫一遍,比如分别试 0.5、0.6、0.7,看被拒绝样本的准确率和拒绝率哪个组合最符合业务诉求。做花卉识别的真实场景里,很多误判其实是“不确定但硬猜”导致的,加了拒识之后整体精度反而更容易看。
我做这类项目养成的习惯是:训练完不急着看训练日志,先把手机里自己拍的几张户外花卉照片放进推理脚本里跑一遍,再对照验证集 acc 判断模型到底行不行。图片集的干净程度永远比网络结构更能决定上线效果。这套流程你已经跟着走了一遍,剩下的就是拿自己的花卉数据集试一次——训练脚本、数据划分、避坑清单都齐了,试一次就会有自己的手感。希望帮到你。
本文还有配套的精品资源,点击获取