news 2026/10/5 7:33:29

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
蝴蝶分类数据集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(f"JSON条目数: {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( weights=sample_weights, num_samples=len(sample_weights), replacement=True ) # 之后在DataLoader中传入sampler参数 # train_loader = DataLoader(train_dataset, batch_size=32, sampler=sampler)

class_weights = 1.0 / class_counts.float()让图片数量少的类别获得更大的权重。WeightedRandomSampler的replacement=True表示允许重复采样,这样每个epoch中稀有类别的图片会被多次抽到,梯度更新更均衡。注意如果你的数据集很小,replacement=True会造成模型对少数图片过拟合,这种情况下可以考虑简单的过采样复制。

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, transform=None): """ 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_size=0.2, stratify=all_labels, random_state=42 ) train_imgs, val_imgs, train_labels, val_labels = train_test_split( train_imgs, train_labels, test_size=0.2, stratify=train_labels, random_state=42 ) print(f"训练集: {len(train_imgs)}, 验证集: {len(val_imgs)}, 测试集: {len(test_imgs)}")

stratify=all_labels强制每个类别在划分后的集合中所占比例和原始数据集一致,这是处理类别不平衡数据集的标准操作。random_state=42固定随机种子,保证每次运行结果一致,方便对比实验。如果你不想引入sklearn,也可以自己写分层划分代码,但原理相同,这里不做展开。

4.3 DataLoader参数配置与epoch设置

有了Dataset,接下来就是DataLoader的配置。很多人在这个环节追求大batch,结果显存爆掉,或者num_workers设成0导致训练速度极慢:

import torch from torch.utils.data import DataLoader # 实例化训练集和验证集的Dataset train_dataset = ButterflyDataset( root_dir='Butterfly20/train', # 已经分割好的训练目录 transform=train_transform ) val_dataset = ButterflyDataset( root_dir='Butterfly20/val', transform=val_transform ) train_loader = DataLoader( train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True # 数据提前加载到GPU内存,减少等待时间 ) val_loader = DataLoader( val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True ) # 估算一个epoch的迭代次数 print(f"训练集batch数: {len(train_loader)}") print(f"验证集batch数: {len(val_loader)}")

pin_memory=True只对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(pretrained=True) # 替换最后一层全连接,改成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}")

pretrained=True会自动下载ImageNet权重到本机缓存,首次运行需要网络连接。named_parameters()遍历所有层的参数名,通过名称判断是否属于layer3、layer4或fc,只有这些层的参数会被训练,其他层保持预训练权重不变。这个策略能防止数据量小时低层特征被破坏,同时加速收敛。

5.2 损失函数与优化器配置

20类分类任务使用交叉熵损失,但要注意类别不平衡时给损失函数加权重。优化器我一般选AdamW,收敛稳定且对学习率不像SGD那样敏感:

import torch.optim as optim # 计算类别权重,解决类别不平衡问题 class_counts = torch.tensor([len(os.listdir(f'Butterfly20/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(weight=class_weights.to(device)) optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)

class_weights = 1.0 / class_counts.float()让样本少的类别获得更大惩罚系数,模型会更重视这些类别的错误。CosineAnnealingLR的学习率变化曲线是从初始值余弦下降到最低再回升,能帮助模型跳出局部最优。T_max=50对应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(f"Epoch [{epoch+1}/{num_epochs}], " f"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, " f"Val 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, dim=1) 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上训练的,推理时也要把输入移动到GPU:input_tensor = input_tensor.to(device),否则会报维度不匹配的错误。

6. 进阶用法与技巧:把这份数据集的价值榨干

训练跑通只是第一步。如果你想把这份20类蝴蝶数据集用到极致,下面这几个进阶方向值得试试,不涉及额外数据,只从现有资源里挖潜力。

6.1 利用genus.txt做层级分类

genus.txt提供了每个物种所属的属名。用这个信息可以构建一个层级分类器:先让模型判断属,再在属内判断具体的种。层级分类的优势在于:即使某两个物种外观极其相似,只要它们的属不同,模型在第一个层级就能区分开。实现上可以训练两个分类器,第一个输出属类别(假设10个属),第二个在属内部做物种分类。这个技巧在处理混淆矩阵中频繁成对错分的类别时非常有效,例如两个物种外形接近,但分属不同属。

6.2 用训练好的模型做特征提取器

如果你不想只做分类,还可以把这个模型当作特征提取器,去掉最后的全连接层,用倒数第二层的输出作为每张图片的特征向量。对于20类数据集,每类平均约50到80张图片,特征向量的维度通常为512(ResNet18的最后一层输出),可以用这些特征向量做聚类、近邻检索或者计算物种间的相似度矩阵。这个方向适合做生物多样性研究的辅助分析,比如观察数据集中哪些物种在特征空间中距离最近。

一个简单做法是记录最后一层卷积的输出:

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(pretrained=True) feature_extractor = FeatureExtractor(backbone) # 之后对每张图片前向传播得到512维特征向量

nn.Sequential(*list(backbone.children())[:-1])把ResNet18的卷积层、池化层全部保留,只去掉最后的全连接层。这样前向传播的输出不是分类概率,而是图片的图像特征。如果你已训练好分类模型,直接用它的权重替换掉pretrained=True的位置,提取的特征就带上了蝴蝶特有的语义信息,效果通常会更好。

6.3 模型集成与置信度校准

数据集规模不大时,单个模型的泛化能力有限。使用三个不同初始化或不同网络结构的模型做集成,通常能带来一到两个百分点的准确率提升。具体做法是分别训练ResNet18、ResNet34和MobileNetV3,推理时对三者的softmax输出取平均:

probs = (probs_model1 + probs_model2 + probs_model3) / 3 _, final_pred = torch.max(probs, dim=1)

probs_model1等是各模型对同一张图片的softmax输出,形状为(1, 20)。取平均值后,三个模型一致同意的类别的概率会被加权放大,意见分歧的类别概率会趋于平滑,能降低误判风险。做过集成的都知道,模型之间差异越大,集成效果越好,所以建议三个模型用不同的数据增强方案训练。

最后说说我的习惯:每次拿到新数据集,不管多急着开训,我都会强制走一遍「校验文件数量 → 检查JSON路径 → 测试单张推理 → 训练一个小epoch验证loss下降」这四个步骤。这套流程看起来很机械,但每次都能在数据准备阶段拦下问题。如果你也遇到过训练半天发现数据不对的情况,不妨试试这个顺序,能省下好几个晚上的调试时间。数据集的坑永远是越早踩越不值钱,希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/5 7:33:03

Cesium结合Heatmap.js实现动态洪水模拟:无需GLSL的轻量方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/5 7:32:38

KeyarchOS性能基线实战:UnixBench完整跑分指南

前阵子给一批新服务器做上线前的性能基线&#xff0c;几台同款机器装的是浪潮信息KeyarchOS&#xff08;KOS&#xff09;。业务同学反馈说某些批处理任务偶尔变慢&#xff0c;我需要先确认问题出在系统配置还是硬件本身&#xff0c;于是第一件事就是把UnixBench拉起来跑了一轮。…

作者头像 李华
网站建设 2026/10/5 7:31:56

单细胞热图改造:从颜色块到多维结构视图的实践指南

前些天整理单细胞项目结果&#xff0c;又被审稿人问了一句“你的热图除了颜色深浅还能看出什么”。这句话戳到我了。单细胞转录组分析里&#xff0c;热图几乎是标配&#xff0c;但绝大多数人画出来的热图&#xff0c;就是一个“表达量颜色块”&#xff0c;既看不出细胞亚群的差…

作者头像 李华
网站建设 2026/10/5 7:31:01

OneOS OTA远程升级全解析:从双区备份到灰度发布

1. 从一次“上门维护”开始&#xff1a;为什么物联网设备必须学会OTA升级做物联网开发的同仁应该都有过这种经历&#xff1a;产品已经铺到现场&#xff0c;甚至铺了几百上千台&#xff0c;突然发现某个固件版本存在一个隐蔽的逻辑bug&#xff0c;或者客户提了一个新需求需要改配…

作者头像 李华
网站建设 2026/10/5 7:31:01

一键化远程桌面启用脚本:Windows与Linux部署及排障指南

远程桌面这东西&#xff0c;属于那种“平时想不起来&#xff0c;真到用的时候急得跳脚”的功能。我最早真正较真去搞&#xff0c;是帮同事远程处理一台异地机房的服务器故障&#xff1a;机器放在现场&#xff0c;人在办公室&#xff0c;偏偏远程桌面没开、服务也没起来&#xf…

作者头像 李华
网站建设 2026/10/5 7:31:00

COCO转YOLO避坑指南:传送带异物检测数据集处理全流程

简介&#xff1a;传送带异物检测识别数据集聚焦工业生产线上的异物检测任务&#xff0c;可支持识别铁棍、垃圾等目标&#xff0c;面向计算机视觉算法工程师、工业质检系统开发者以及相关专业学生&#xff0c;能够为算法验证、项目演示和课程实验提供带标注的真实场景数据。数据…

作者头像 李华