ARTICLE DETAIL

资讯详情

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

蝴蝶分类数据集20类:用PyTorch跑通图像分类全流程

蝴蝶分类数据集20类:用PyTorch跑通图像分类全流程 简介蝴蝶分类数据集包含20个常见蝴蝶物种类别适用于图像识别、深度学习模型训练、生物多样性研究及教学演示等场景可帮助研究者与学习者快速获得带标注的图像样本。整个压缩包共1870个文件大小约60.96MB主体为1866张JPG格式图片并配有1个JSON字典和2个TXT文本文件。JSON文件中存储了每张图片的路径、物种名、属名等元数据TXT文件分别列举全部物种名与对应属名便于生成分类标签和进行数据检索。图像按类别放在独立文件夹中每类提供多张不同角度或状态的样本为训练卷积神经网络等模型提供了较全面的视角变化。目前已有119人学习下载适合机器学习初学者在分类任务中直接使用也可供生物学者分析物种分布与进化关系。数据集结构简洁、标注完整省去了手动整理标签的流程可即下即用。1. 蝴蝶分类数据集20类一个能直接训练的图像分类资源做图像分类的同行应该都遇到过这种尴尬想验证一个新网络结构或者调试一套训练管线但手里能用的数据集要么太大要下半天要么类别乱七八糟没法对齐实验。这个名为「蝴蝶分类数据集20类」的zip资源解压后就是一个立即可用的图像分类数据包20个类别每个类别对应一种蝴蝶物种图片按文件夹组织好还带着Butterfly20_dict.json、species.txt、genus.txt三个标注文件。这意味着你不需要自己写爬虫抓图、不用手动清理坏图直接拿来就能跑通一个完整的分类训练流程。适合两类人刚接触图像分类、想拿一份干净数据练手的初学者以及需要快速验证模型改动效果、不想在数据准备上花时间的工程研究者。它的价值不在于数据量大而在于「元数据齐全、目录结构规整、压缩包打开即用」这三点能把你从数据清理的黑洞里拉出来直接进训练环节。2. 数据组织与元数据解析先搞清压缩包里到底有什么拿到任何数据集第一步都不是开训练而是把目录结构和标注文件摸清楚。很多翻车现场都是因为数据读取阶段就出了错后面所有训练结果全报废。这一章把压缩包内的文件布局和你需要关注的字段全部拆开讲。2.1 Butterfly20目录布局与图片命名规则压缩包解压后核心目录是Butterfly20里面按类别分子目录。每个子目录名通常就是物种标识目录下放置该物种的多张图片。图片命名是纯数字加序号比如019.jpg、050.jpg、077.jpg、126.jpg没有中文也没有空格这点对Linux环境和Windows环境都很友好不会因为编码问题导致读取失败。我一般拿到第一件事是确认每个类别的图片数量是否均匀。用一段bash脚本扫一下# 统计每个子目录下的图片数量并打印到终端 for dir in Butterfly20/*/; do count$(ls $dir | wc -l) echo $dir: $count done这段脚本遍历Butterfly20下所有子目录统计每个目录的文件数。wc -l统计行数这里等价于文件数量。跑完后你会对数据分布有个直观认识。如果某个类别的图片数量明显少比如个位数后续训练要考虑数据增强或过采样来平衡。图片格式是常见的jpg尺寸没有做统一裁剪这意味着原始图片可能有大有小。做训练时不要直接resize到一个固定尺寸然后开训而是应该在数据加载器里统一处理。常见做法是短边缩放加中心裁剪或者直接resize到统一尺寸两种方案在第三章会给出具体代码。2.2 species.txt与genus.txt从标签到生物学层级species.txt和genus.txt是这份数据集和普通图片文件夹最大的区别。species.txt每行一个物种名共20行顺序就是类别索引的顺序。genus.txt每行一个属名同一属下的物种可能共享一行前缀信息。这两个文件的意义在于分类任务的标签不只是「类别0到类别19」而是有明确生物学语义的物种名和属名。训练时类别标签和species.txt的对应关系是类别索引0对应第一行物种名索引1对应第二行以此类推。很多人在做推理时要输出中文物种名但训练时用的是英文这里就需要一张映射表。我习惯在训练前先加载这两个文件动态生成标签映射而不是硬编码在代码里# 读取物种名和属名生成标签映射字典 with open(species.txt, r) as f: species_list [line.strip() for line in f.readlines()] with open(genus.txt, r) as f: genus_list [line.strip() for line in f.readlines()] # 类别索引 - 物种全名 - 属名三层映射 idx_to_species {i: name for i, name in enumerate(species_list)} species_to_genus {species: genus for species, genus in zip(species_list, genus_list)} idx_to_genus {i: species_to_genus[species_list[i]] for i in range(len(species_list))} print(类别数量:, len(species_list)) print(索引0对应的物种:, idx_to_species[0])这段代码把20个类别做成三个字典idx_to_species用于训练输出时显示物种名species_to_genus用于从物种查属idx_to_genus可以直接输出预测的属级别结果。zip函数将两个列表配对注意species.txt和genus.txt的行数必须一一对应否则zip会静默截断到较短的长度这是老手也会踩的坑。2.3 Butterfly20_dict.json路径、元数据与模型训练的衔接Butterfly20_dict.json是这份数据集最值得留意的文件。它通常是一个字典结构键是图片路径或图片文件名值包含物种名、属名和其他可能的描述性字段。这个JSON文件在训练管线里的用处主要有三个一是做数据集划分train/val/test二是做类别统计三是排查图片和标注是否匹配。我一般会先把这个JSON读进来检查键的数量和Butterfly20目录里实际图片数是否一致import json with open(Butterfly20_dict.json, r) as f: data json.load(f) print(JSON中图片条目数:, len(data)) # 展示前两条记录确认字段结构 for i, (img_path, info) in enumerate(data.items()): if i 2: print(图片:, img_path) print(标注:, info) else: break这段代码输出JSON条目的总数和前两条记录的结构。注意json.load返回的是字典data.items()遍历所有的键值对。这里有个关键点通常JSON里记录的图片路径是相对路径比如Butterfly20/species_name/019.jpg但解压后你的实际目录前缀可能不同比如你解压到了/home/user/data/下这时候直接用JSON里的路径会报文件不存在。解决方法是写一个路径前缀修正函数在读取JSON后统一替换路径前缀。常见做法是提取JSON中所有图片路径的公共前缀然后替换成你本地的实际根目录import os from pathlib import Path # JSON里记录的图片路径 sample_path list(data.keys())[0] print(JSON中的路径示例:, sample_path) # 你本地的实际数据根目录 local_root Path(./) # 修正函数提取相对路径去掉公共目录前缀 def fix_path(json_path): # 只保留 filename 最后两级类别名/图片名 parts Path(json_path).parts rel_path os.path.join(*parts[-2:]) return str(local_root / Butterfly20 / rel_path) fixed_path fix_path(sample_path) print(修正后路径:, fixed_path)这段代码不依赖JSON里可能存在的固定前缀而是直接提取路径的最后两部分重新拼接到你的本地目录下。Path(json_path).parts会把路径拆成元组比如(Butterfly20, some_species, 019.jpg)取最后两个元素就是(some_species, 019.jpg)再和本地根目录拼接。这种写法能兼容绝大多数路径不一致的情况前提是JSON里记录的路径保证最后两级是「类别名/图片名」。3. 数据校验与常见问题排查训练前必做的四道检查这一章是给所有急着开训的人泼冷水。数据集的完整度和标注质量直接决定模型效果作者提供的文件清单能确认大部分信息但解压后你手里的文件是否和JSON标注一致、图片是否有损坏、类别是否失衡这些都必须自己验证。我见过太多人跳过这步直接训练最后loss不降反升还不知道去哪排查。3.1 图片文件与JSON标注的数量一致性检查先做一个最基础的校验JSON条目数和实际图片文件数是否一致。不一致的原因很多可能是压缩包本身少了文件也可能是解压过程中出现了路径错误、文件名被截断。用一段Python脚本做全量校验import os import json with open(Butterfly20_dict.json, r) as f: data json.load(f) # 统计JSON条目数 json_count len(data) print(fJSON条目数: {json_count}) # 统计实际图片数 actual_images [] for root, dirs, files in os.walk(Butterfly20): for file in files: if file.endswith((.jpg, .jpeg, .png)): actual_images.append(os.path.join(root, file)) print(f实际图片数: {len(actual_images)}) # 找出JSON中路径对应的文件是否存在 missing [] for img_path in data.keys(): # 取相对于Butterfly20的路径 parts img_path.split(/) # 尝试定位到本地 candidate os.path.join(Butterfly20, *parts[-2:]) if not os.path.exists(candidate): missing.append(img_path) print(f缺失文件数: {len(missing)}) if missing: print(样例:, missing[:5])这段脚本做了双向校验先统计实际图片数量再逐个检查JSON中记录的每张图片是否存在。os.walk遍历所有子目录统计jpg、jpeg、png文件split(/)和parts[-2:]处理的是不同操作系统下路径分隔符不统一的问题。如果你发现缺失数量很多不要试图用代码自动补齐先查找压缩包是否漏解压或者JSON本身有问题。3.2 图片完整性验证警惕“幽灵文件”有的图片文件虽然存在但文件头损坏读取时直接报错。这种情况在深度学习中很常见特别是从网络爬取的数据集。验证方式很简单用PIL尝试打开所有图片如果某张图打不开或者格式不匹配就把它标记出来from PIL import Image import os broken_images [] count 0 for root, dirs, files in os.walk(Butterfly20): for file in files: if file.endswith((.jpg, .jpeg, .png)): img_path os.path.join(root, file) count 1 try: with Image.open(img_path) as img: img.verify() # 只校验文件头不加载全图速度快 except Exception as e: broken_images.append((img_path, str(e))) print(f检查图片总数: {count}) print(f损坏图片数: {len(broken_images)}) if broken_images: for path, err in broken_images[:10]: print(f损坏: {path}, 错误: {err})Image.verify()和Image.open()的区别在于verify()只读取文件头确认格式合法不会把整张图片解码到内存所以批量检查时效率很高。如果你发现损坏图片数量很少个位数直接删除或者用同一类别的其他图片替代即可如果数量很大说明数据源本身有问题不建议直接训练。3.3 常见坑一类别标签与目录名不匹配导致训练精度崩盘现象训练集准确率很高但验证集准确率始终在10%接近随机猜测上下徘徊检查数据加载代码也没发现明显错误。原因species.txt中物种名的顺序和Butterfly20下子目录的排序不一致。很多人会用os.listdir()或glob直接读取目录列表作为标签顺序而文件系统不保证目录的字母顺序和species.txt的行顺序一致。这会导致你的模型实际上一直在用「甲标签」学「乙图片」验证时自然全面崩盘。解决强制用species.txt作为唯一标签顺序来源构造目录名到类别索引的显式映射import os # 读取species.txt作为标签顺序的唯一标准 with open(species.txt, r) as f: species_list [line.strip() for line in f.readlines()] # 建立目录名到类别索引的映射 dir_to_idx {} for idx, species in enumerate(species_list): dir_filename species.replace( , _) # 物种名可能带空格目录里可能用下划线 dir_to_idx[dir_filename] idx print(f映射条目数: {len(dir_to_idx)}) print(样例:, list(dir_to_idx.items())[:3]) # 遍历目录时用这个映射而不是枚举目录顺序 for dirname in os.listdir(Butterfly20): if dirname in dir_to_idx: print(f{dirname} - 类别索引 {dir_to_idx[dirname]}) else: print(f警告: 目录 {dirname} 不在species.txt中)这段代码把物种名和目录名做了一次显式关联。replace( , _)是为了处理物种名带空格的情况——部分数据集的目录会用下划线替代空格。跑完后你会立刻发现哪些目录和标签对不上在训练前就修正掉。3.4 常见坑二直接resize导致蝴蝶特征被压缩变形现象训练loss能正常下降模型收敛速度也正常但推理时对真实图片的分类准确率很差。原因数据集里的蝴蝶图片本身是自然拍摄的照片蝴蝶在画面中的占比和位置各不相同。如果直接resize到224x224小尺寸图片里的蝴蝶可能被压到10像素大小特征完全丢失。解决在数据加载时先做短边等比缩放再做中心裁剪到目标尺寸。PyTorch里用transforms.Resize加transforms.CenterCrop组合from torchvision import transforms # 训练集使用随机裁剪和水平翻转增强泛化能力 train_transform transforms.Compose([ transforms.Resize(256), # 短边缩放到256 transforms.RandomResizedCrop(224), # 随机裁剪到224带缩放变化 transforms.RandomHorizontalFlip(), # 随机水平翻转 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 验证集使用中心裁剪固定预处理流程 val_transform 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]) ])Resize(256)先把短边缩放到256像素RandomResizedCrop(224)在缩放后的图上随机选取区域并裁剪到224这相当于同时完成了尺度增强和位置增强模型能学到不同大小和位置的蝴蝶特征。Normalize里的均值和标准差是ImageNet的标准值如果你用的是在ImageNet上预训练的模型必须保持一致如果是自己从头训练这两个值可以缺省。3.5 常见坑三类别不平衡导致小类别永远分不准现象总体准确率有90%但看混淆矩阵时发现某一两个类别的召回率极低几乎全被分到其他类别。原因自然拍摄的蝴蝶数据集中常见物种的图片数量可能是稀有物种的十倍以上。模型在训练时倾向于把大量样本的类别学得更好稀有类别因为样本太少梯度更新不足。解决先从统计入手确认类别分布情况再考虑用权重或重采样import os from collections import Counter # 统计每个类别子目录的图片数量 category_counts Counter() for root, dirs, files in os.walk(Butterfly20): if root ! Butterfly20: category os.path.basename(root) category_counts[category] len([f for f in files if f.endswith((.jpg, .jpeg, .png))]) print(各类别图片数量:) for cat, count in category_counts.most_common(): print(f {cat}: {count}) # 计算不平衡率 if category_counts: max_count max(category_counts.values()) min_count min(category_counts.values()) print(f\n最大类别数: {max_count}, 最小类别数: {min_count}) print(f不平衡率: {max_count / min_count:.2f}:1)这段代码用Counter统计每个子目录的图片数量并给出最大/最小类别的不平衡率。如果比例超过5:1就建议在训练中加入加权采样。PyTorch里可以用WeightedRandomSampler实现import torch from torch.utils.data import WeightedRandomSampler # 假设你已经构建了dataset和标签列表labels # 计算每个类别的权重稀有类别权重更高 class_counts torch.bincount(torch.tensor(labels)) class_weights 1.0 / class_counts.float() sample_weights class_weights[labels] # 每个样本的权重 sampler WeightedRandomSampler( weightssample_weights, num_sampleslen(sample_weights), replacementTrue ) # 之后在DataLoader中传入sampler参数 # train_loader DataLoader(train_dataset, batch_size32, samplersampler)class_weights 1.0 / class_counts.float()让图片数量少的类别获得更大的权重。WeightedRandomSampler的replacementTrue表示允许重复采样这样每个epoch中稀有类别的图片会被多次抽到梯度更新更均衡。注意如果你的数据集很小replacementTrue会造成模型对少数图片过拟合这种情况下可以考虑简单的过采样复制。4. 训练前的数据划分与加载写一份你自己的数据集类数据集校验通过后下一步就是把Butterfly20改造成一个PyTorch可以直接喂给模型的Dataset。我之前见不少新手直接手写循环读取图片然后拼成numpy数组喂给模型这种方式在20类这种小规模数据集上勉强能跑但完全没有扩展性一旦要加验证集划分、做数据增强代码就要推翻重来。这一章从Dataset类的写法开始到划分策略一次性讲完。4.1 自定义Dataset直接读取目录结构和JSON标注最稳妥的做法是以Butterfly20目录结构为准写Dataset利用子目录名作为类别标签同时交叉验证species.txt的顺序import os from PIL import Image import torch from torch.utils.data import Dataset from torchvision import transforms class ButterflyDataset(Dataset): def __init__(self, root_dir, transformNone): root_dir: Butterfly20目录的绝对或相对路径 transform: 预处理操作组合通常来自torchvision.transforms self.root_dir root_dir self.transform transform self.image_paths [] self.labels [] self.class_names sorted(os.listdir(root_dir)) # 建议用sorted保证顺序稳定 self.class_to_idx {name: idx for idx, name in enumerate(self.class_names)} # 遍历所有子目录 for class_name in self.class_names: class_dir os.path.join(root_dir, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): if img_name.endswith((.jpg, .jpeg, .png)): self.image_paths.append(os.path.join(class_dir, img_name)) self.labels.append(self.class_to_idx[class_name]) def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img_path self.image_paths[idx] image Image.open(img_path).convert(RGB) # 统一转成RGB防止灰度图报错 label self.labels[idx] if self.transform: image self.transform(image) return image, label这里有几个关键设计convert(RGB)强制把所有图片统一成三通道RGB避免数据集里混有灰度图或RGBA图导致训练时维度不一致的错误sorted(os.listdir())确保类别顺序稳定配合species.txt使用时建议手动验证一下顺序是否一致。4.2 数据集划分不只在初始化里做随机切分严格的做法是先把所有样本按照类别分层划分成训练集、验证集和测试集然后各自实例化Dataset。分层划分的意义在于保证每个类别在三个集合中的比例一致import random import os from sklearn.model_selection import train_test_split # 获取所有样本路径和标签 all_images [] all_labels [] for class_name in sorted(os.listdir(Butterfly20)): class_dir os.path.join(Butterfly20, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): if img_name.endswith((.jpg, .jpeg, .png)): all_images.append(os.path.join(class_dir, img_name)) all_labels.append(class_name) # 先分出测试集再在剩余数据中分验证集 train_imgs, test_imgs, train_labels, test_labels train_test_split( all_images, all_labels, test_size0.2, stratifyall_labels, random_state42 ) train_imgs, val_imgs, train_labels, val_labels train_test_split( train_imgs, train_labels, test_size0.2, stratifytrain_labels, random_state42 ) print(f训练集: {len(train_imgs)}, 验证集: {len(val_imgs)}, 测试集: {len(test_imgs)})stratifyall_labels强制每个类别在划分后的集合中所占比例和原始数据集一致这是处理类别不平衡数据集的标准操作。random_state42固定随机种子保证每次运行结果一致方便对比实验。如果你不想引入sklearn也可以自己写分层划分代码但原理相同这里不做展开。4.3 DataLoader参数配置与epoch设置有了Dataset接下来就是DataLoader的配置。很多人在这个环节追求大batch结果显存爆掉或者num_workers设成0导致训练速度极慢import torch from torch.utils.data import DataLoader # 实例化训练集和验证集的Dataset train_dataset ButterflyDataset( root_dirButterfly20/train, # 已经分割好的训练目录 transformtrain_transform ) val_dataset ButterflyDataset( root_dirButterfly20/val, transformval_transform ) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue # 数据提前加载到GPU内存减少等待时间 ) val_loader DataLoader( val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue ) # 估算一个epoch的迭代次数 print(f训练集batch数: {len(train_loader)}) print(f验证集batch数: {len(val_loader)})pin_memoryTrue只对GPU训练有意义它让数据在CPU端分配时使用页锁定内存数据传输到GPU时更快。num_workers建议设置为CPU核心数的一半到三分之二太高会导致进程间通信开销大于实际加载时间。对于20类蝴蝶这种规模的数据集epoch数设置50到100足够过多反而可能导致过拟合。5. 训练与微调用ResNet跑通20类蝴蝶分类数据管线就绪后剩下的就是模型选择、训练参数配置和收敛判断。20类分类任务属于轻量级视觉任务不需要上几十上百层的超大网络用ResNet18或ResNet34这类中等规模网络就能取得不错的效果。如果你有一块入门级GPU显存4GB以上这个任务的显存占用非常小。5.1 模型选择与迁移学习策略蝴蝶分类的图片和ImageNet的自然图像有共通的低层特征边缘、纹理、颜色所以强烈建议使用ImageNet预训练权重做迁移学习。常见做法是冻结预训练模型的前几层只微调后面几层和全连接层import torchvision.models as models import torch.nn as nn # 加载预训练的ResNet18 model models.resnet18(pretrainedTrue) # 替换最后一层全连接改成20类输出 num_features model.fc.in_features model.fc nn.Linear(num_features, 20) # 冻结前几层参数只训练layer3、layer4和最后的fc层 for name, param in model.named_parameters(): if layer3 not in name and layer4 not in name and fc not in name: param.requires_grad False # 打印参数量确认哪些层参与训练 trainable_params sum(p.numel() for p in model.parameters() if p.requires_grad) total_params sum(p.numel() for p in model.parameters()) print(f可训练参数: {trainable_params}, 总参数: {total_params})pretrainedTrue会自动下载ImageNet权重到本机缓存首次运行需要网络连接。named_parameters()遍历所有层的参数名通过名称判断是否属于layer3、layer4或fc只有这些层的参数会被训练其他层保持预训练权重不变。这个策略能防止数据量小时低层特征被破坏同时加速收敛。5.2 损失函数与优化器配置20类分类任务使用交叉熵损失但要注意类别不平衡时给损失函数加权重。优化器我一般选AdamW收敛稳定且对学习率不像SGD那样敏感import torch.optim as optim # 计算类别权重解决类别不平衡问题 class_counts torch.tensor([len(os.listdir(fButterfly20/train/{cls})) for cls in sorted(os.listdir(Butterfly20/train))]) class_weights 1.0 / class_counts.float() # 归一化权重让平均值接近1 class_weights class_weights / class_weights.mean() criterion nn.CrossEntropyLoss(weightclass_weights.to(device)) optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50)class_weights 1.0 / class_counts.float()让样本少的类别获得更大惩罚系数模型会更重视这些类别的错误。CosineAnnealingLR的学习率变化曲线是从初始值余弦下降到最低再回升能帮助模型跳出局部最优。T_max50对应50个epoch。5.3 训练循环调试与loss监控训练循环本身不复杂但要注意几个容易出错的位置model.train()和model.eval()的切换、梯度清零的时机、验证阶段不要计算梯度。完整循环如下device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) num_epochs 50 best_val_acc 0.0 for epoch in range(num_epochs): # 训练阶段 model.train() running_loss 0.0 correct 0 total 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() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() train_acc 100.0 * correct / total train_loss running_loss / len(train_loader) # 验证阶段 model.eval() val_correct 0 val_total 0 with torch.no_grad(): # 不计算梯度节省显存 for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) val_total labels.size(0) val_correct (predicted labels).sum().item() val_acc 100.0 * val_correct / val_total print(fEpoch [{epoch1}/{num_epochs}], fTrain Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, fVal Acc: {val_acc:.2f}%) # 保存验证集准确率最高的模型 if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_butterfly_model.pth)optimizer.zero_grad()必须在每个batch开始前调用否则梯度会累积等同于增大batch_size导致模型震荡。torch.no_grad()在验证阶段省去了梯度计算和存储的开销能明显减少显存占用。判断是否过拟合的直观信号是训练集准确率继续上升但验证集准确率停滞或下降此时应该考虑加早停或减小学习率。5.4 推理验证把预测结果映射回物种名训练完成后你要把模型的数字类别输出映射回species.txt里的物种名。这一步容易出的问题是类别索引和物种名对不上原因通常是Dataset内部用sorted(os.listdir())排序而排序规则和species.txt的行序不一致import torch from PIL import Image from torchvision import transforms def predict_image(model, img_path, idx_to_species): model: 训练好的模型 img_path: 待推理图片路径 idx_to_species: 索引到物种名的映射字典 model.eval() image Image.open(img_path).convert(RGB) transform 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]) ]) input_tensor transform(image).unsqueeze(0) # 增加batch维度 with torch.no_grad(): outputs model(input_tensor) probabilities torch.softmax(outputs, dim1) top_prob, top_idx torch.topk(probabilities, 3) # 取概率最高的3个 print(f预测结果:) for i in range(top_prob.size(1)): idx top_idx[0, i].item() prob top_prob[0, i].item() print(f {idx_to_species[idx]}: {prob * 100:.2f}%)torch.topk返回概率最大的前k个索引和对应的概率值。这里有一个关键点unsqueeze(0)把形状从(3, 224, 224)扩展成(1, 3, 224, 224)满足模型输入要求。如果模型是在GPU上训练的推理时也要把输入移动到GPUinput_tensor input_tensor.to(device)否则会报维度不匹配的错误。6. 进阶用法与技巧把这份数据集的价值榨干训练跑通只是第一步。如果你想把这份20类蝴蝶数据集用到极致下面这几个进阶方向值得试试不涉及额外数据只从现有资源里挖潜力。6.1 利用genus.txt做层级分类genus.txt提供了每个物种所属的属名。用这个信息可以构建一个层级分类器先让模型判断属再在属内判断具体的种。层级分类的优势在于即使某两个物种外观极其相似只要它们的属不同模型在第一个层级就能区分开。实现上可以训练两个分类器第一个输出属类别假设10个属第二个在属内部做物种分类。这个技巧在处理混淆矩阵中频繁成对错分的类别时非常有效例如两个物种外形接近但分属不同属。6.2 用训练好的模型做特征提取器如果你不想只做分类还可以把这个模型当作特征提取器去掉最后的全连接层用倒数第二层的输出作为每张图片的特征向量。对于20类数据集每类平均约50到80张图片特征向量的维度通常为512ResNet18的最后一层输出可以用这些特征向量做聚类、近邻检索或者计算物种间的相似度矩阵。这个方向适合做生物多样性研究的辅助分析比如观察数据集中哪些物种在特征空间中距离最近。一个简单做法是记录最后一层卷积的输出import torch import torch.nn as nn # 提取ResNet18倒数第二层特征 class FeatureExtractor(nn.Module): def __init__(self, backbone): super().__init__() # 去掉最后一层全连接 self.features nn.Sequential(*list(backbone.children())[:-1]) def forward(self, x): x self.features(x) return x.view(x.size(0), -1) # 展平成向量 backbone models.resnet18(pretrainedTrue) feature_extractor FeatureExtractor(backbone) # 之后对每张图片前向传播得到512维特征向量nn.Sequential(*list(backbone.children())[:-1])把ResNet18的卷积层、池化层全部保留只去掉最后的全连接层。这样前向传播的输出不是分类概率而是图片的图像特征。如果你已训练好分类模型直接用它的权重替换掉pretrainedTrue的位置提取的特征就带上了蝴蝶特有的语义信息效果通常会更好。6.3 模型集成与置信度校准数据集规模不大时单个模型的泛化能力有限。使用三个不同初始化或不同网络结构的模型做集成通常能带来一到两个百分点的准确率提升。具体做法是分别训练ResNet18、ResNet34和MobileNetV3推理时对三者的softmax输出取平均probs (probs_model1 probs_model2 probs_model3) / 3 _, final_pred torch.max(probs, dim1)probs_model1等是各模型对同一张图片的softmax输出形状为(1, 20)。取平均值后三个模型一致同意的类别的概率会被加权放大意见分歧的类别概率会趋于平滑能降低误判风险。做过集成的都知道模型之间差异越大集成效果越好所以建议三个模型用不同的数据增强方案训练。最后说说我的习惯每次拿到新数据集不管多急着开训我都会强制走一遍「校验文件数量 → 检查JSON路径 → 测试单张推理 → 训练一个小epoch验证loss下降」这四个步骤。这套流程看起来很机械但每次都能在数据准备阶段拦下问题。如果你也遇到过训练半天发现数据不对的情况不妨试试这个顺序能省下好几个晚上的调试时间。数据集的坑永远是越早踩越不值钱希望帮到你。本文还有配套的精品资源点击获取
返回列表