简介:水果分类数据集压缩包面向机器学习、计算机视觉与数据挖掘初学者,内含苹果、香蕉、葡萄、橙子、梨五类常见水果的标注图片,可用于图像分类模型的训练与评估,也可作为理解分类任务和标签组织方式的入门素材。压缩包共1310个文件,主体为1306张JPG图片,另附2个标签列表、1个JSON配置及1个Python脚本,分别用于梳理类别清单、记录映射关系和辅助数据加载,整体大小约14.07MB,轻量易用。目前已有3625人学习下载,是实践监督学习流程的热门资源。借助这份数据集,使用者不仅能获得带标签的水果样本,还能通过附带的脚本快速完成数据划分、预处理和特征提取,进而尝试CNN等模型,经历从数据准备到分类效果评估的完整实战过程。
1. 水果分类数据集:这份 rar 里到底装了什么,值不值得花时间解开
做图像分类的入门项目,十个有八个会从水果数据集开始。原因很朴素:类别数量适中、图片特征明显、背景相对干净,拿来做迁移学习或者从零训练一个小网络,都能在不太长的训练时间内看到损失下降。但真正拿到一份「水果分类数据集 fruits分类数据集.rar」,第一步往往不是写训练脚本,而是先跟这个压缩包搏斗——它是什么格式、里面目录怎么排、标签是文件夹名还是单独的 CSV、图片尺寸统不统一,这些信息在解压之前全是黑匣子。这篇笔记就是把从拿到 rar 到跑通第一个训练脚本的完整路径拆开讲,包含解压工具的选择、目录结构的检查方法、数据集划分的注意事项,以及我实际踩过的几个坑。适合刚接触图像分类、手里正好有一份 rar 格式数据集、又不想在数据准备阶段耗掉整个下午的读者。
2. 解压 .rar 数据集:工具选型与两个高频翻车点
2.1 为什么 rar 格式的数据集还这么常见,以及解压工具怎么选
深度学习数据集大多以 zip 或 tar.gz 分发,但 rar 依然活跃在网盘分享和学术资源的私下流传里,原因是它的压缩率在相同配置下通常比 zip 高几个百分点,对动辄几个 GB 的图片集来说,能省下不少上传时间和网盘空间。代价就是生态支持差——Windows 自带资源管理器不认识 rar,Linux 默认也没有解 rar 的命令行工具,macOS 的归档实用工具同样无能为力。
常见做法是装一个跨平台的解压工具。Windows 上 7-Zip 是最稳妥的选择,免费、开源、无广告,右键菜单直接有「解压到当前文件夹」;macOS 上可以用 The Unarchiver;Linux 服务器上则是安装 p7zip 系列包。这里要单独提一句:网上搜索 rar 解压软件时,很容易下到带广告推广的「万能压缩」类软件,界面花哨但解压速度慢,还会在后台弹推广。我一般直接固定用 7-Zip,它同时支持解压 rar、zip、7z,生成环境里用命令行版本,干净利落。
2.2 Linux 和 Windows 下的实际操作命令
如果你在本地 Windows 上操作,图形界面右键解压就够了,但建议养成用命令行解压的习惯,尤其是在处理大批量数据时——命令行能保留完整路径、避免部分文件解压失败时弹窗中断,也方便写入脚本做自动化。
Windows 下安装了 7-Zip 后,打开 PowerShell 或 CMD,用 7z 命令解压:
7z x fruits分类数据集.rar -oD:\datasets\fruits -yx表示保留压缩包内的目录结构完整解压,-o后面紧跟解压目标路径,注意-o和路径之间不能有空格,-y是遇到同名文件直接覆盖。如果你用的是 Linux 服务器:
sudo apt install p7zip-full 7z x fruits分类数据集.rar -o/home/user/datasets/fruits -yp7zip-full是 7-Zip 在 Linux 下的实现,装好之后7z命令就可用。如果服务器上没有 root 权限,也可以用unar(The Unarchiver 的命令行版),它对 rar 的处理同样可靠。解压完成后先别急着看图片,先执行find . -type f | wc -l统计一下文件数量,再du -sh看一下总大小,这两个数字能帮你快速判断解压是否完整,防止中途断电或磁盘空间不足导致静默失败。
2.3 解压密码和文件名乱码的处理经验
部分分享者会给数据集压缩包加密码,解压时遇到提示输密码,先检查下载页面或分享说明里有没有附带密码。常见的套路是「解压密码在文件名后缀」或「关注公众号获取」,这类其实都能在下载页面找到。如果压缩包是加密的而你又完全不知道密码,网上所谓的 rar 密码破解工具大多不可靠——暴力破解 rar 密码的时间成本极高,8 位混合密码在普通 PC 上可能要跑几百年。我的建议是换个下载源,别在解密上花时间。有工具叫 rar password cracker,原理是字典攻击,对弱密码偶尔有效,试一两次可以,别指望它是万能钥匙。
文件名乱码是第二个高频问题,多见于国内分享的资源。rar 在 Windows 下用 GBK 编码文件名,在 Linux 或 macOS 下解压时被当成 UTF-8 读取,于是出现「Ê¥Â」这类乱码目录名。解决方式是在 Linux 下用convmv或7z -mcp=936指定编码:
7z x fruits分类数据集.rar -mcp=936-mcp=936表示按 GBK 读取文件名,解压出来目录名就正常了,Windows 本地解压一般不需要加这个参数因为系统默认就是 GBK。
提示:解压完成后先进入最外层目录用
ls看一眼目录结构。如果发现目录名是一串乱码,立刻删掉重新用-mcp=936解压,不要手动一个个重命名,几十个类别文件夹手工改名的成本你承受不起。
3. 数据集内部结构:类别目录、标注格式与文件清单核对
3.1 三种常见水果数据集的目录组织方式
解压完成之后,接下来要面对的是「数据集到底长什么样」的问题。我见过的大多数水果分类 rar 包,内部结构无非是以下三种之一。
第一种是标准的 ImageFolder 结构:外层是一个主目录,里面每个类别一个子文件夹,子文件夹名就是类别标签,图片直接放在子文件夹里。例如apple/、banana/、orange/,每个文件夹内是若干张 jpg。这种结构最简单,PyTorch 的torchvision.datasets.ImageFolder直接就能读取,标签按文件夹名称的字母顺序自动映射为数字。
第二种是图片全部堆在一个文件夹里,旁边配一个labels.csv或train.csv,CSV 里两列:文件名和对应的类别名。这种结构需要自己写代码把图片路径和标签对应起来,用 pandas 读 CSV 再做映射。
第三种是混合型:训练集和测试集已经分好,train/下是类别子文件夹,test/下也是类别子文件夹,但测试集的文件夹里可能没有标签(用于提交结果)或者有部分标签。还有的会额外附一个labels.txt或README.txt说明类别列表。
拿到手先判断属于哪一种,决定了后面所有数据处理代码的写法。不要上来就写训练脚本,先花两分钟看清楚目录结构。
3.2 用脚本做文件完整性检查:损坏图片、空文件夹与类别不平衡
看清结构后,直接写一个 Python 脚本来做全面体检。这个脚本做的事情是:遍历所有图片文件、检查能否被 OpenCV 或 PIL 正常打开、统计每个类别的图片数量、找出损坏文件和非图片文件。
import os from PIL import Image from collections import Counter dataset_root = "fruits_dataset" # 常见图片扩展名 image_exts = {".jpg", ".jpeg", ".png", ".bmp", ".webp"} label_counter = Counter() corrupted_files = [] non_image_files = [] for root, dirs, files in os.walk(dataset_root): for fname in files: fpath = os.path.join(root, fname) ext = os.path.splitext(fname)[1].lower() if ext not in image_exts: non_image_files.append(fpath) continue label = os.path.basename(root) label_counter[label] += 1 # 尝试打开图片,判断是否损坏 try: with Image.open(fpath) as img: img.verify() except Exception as e: corrupted_files.append((fpath, str(e))) print("类别分布:") for label, count in label_counter.most_common(): print(f" {label}: {count}") print(f"损坏图片数: {len(corrupted_files)}") for fpath, err in corrupted_files[:10]: print(f" {fpath} -> {err}") print(f"非图片文件数: {len(non_image_files)}") for fpath in non_image_files[:10]: print(f" {fpath}")这段脚本里,os.walk递归遍历所有子目录,Image.verify()是 PIL 里比较轻量的图片校验方法,只检查文件头和数据完整性,不会把整个图片解码进内存,所以速度很快,几百 MB 的图片集一分钟内能扫完。Counter统计类别分布,方便你一眼看出有没有类别严重不平衡——比如 apple 有 2000 张图而 strawberry 只有 80 张,这种差距训练出来的模型对 apple 严重过拟合。
如果发现损坏图片数量比较多(比如超过总量的 1%),建议直接从数据集中剔除,不要想着靠数据增强补回来。一个损坏图片出现在训练集里会导致训练过程出现莫名其妙的 loss 尖刺,出现在验证集里会导致准确率计算偏差。用一个简单的filter_corrupted.py把损坏文件移到corrupted_backup/目录,比改训练代码更省事。
3.3 类别标签的形式与映射规则
类别标签可能是中文(苹果、香蕉)、英文小写(apple)、带下划线的变体(green_apple)或者 Unnamed 编号(class_0、class_1)。这里有一个重要的原则:标签字符串本身不要当作模型输入的一部分,模型只认识数字索引,所以需要一个稳定的映射字典。
import os from collections import OrderedDict labels = sorted([d for d in os.listdir("fruits_dataset") if os.path.isdir(os.path.join("fruits_dataset", d))]) label_to_idx = {label: idx for idx, label in enumerate(labels)} idx_to_label = {idx: label for label, idx in label_to_idx.items()} print("标签映射:") for label, idx in label_to_idx.items(): print(f" {idx} -> {label}") # 保存映射到 json,供训练和推理时使用 import json with open("label_map.json", "w", encoding="utf-8") as f: json.dump({"label_to_idx": label_to_idx, "idx_to_label": idx_to_label}, f, indent=2)sorted()确保映射的顺序稳定,不会因为文件系统的遍历顺序不同导致每次跑脚本得到的索引不一样。json文件保存一份映射,训练时用它把类别转成数字标签,推理时把模型输出的数字还原成可读的类别名。这个看似简单的操作在后面做模型部署时非常有用——你总不希望推理程序里硬编码一个「0 是苹果,1 是香蕉」的列表。
注意:不要用
os.listdir的默认顺序来建映射,它在不同操作系统上返回的顺序不一致,会导致同样的训练数据在 Windows 上训练和在 Linux 上训练得到完全不同的标签编号。
4. 从原始图片到训练管线:数据划分、预处理与基准训练
4.1 训练集、验证集、测试集的划分比例与方法
结构检查完、损坏文件清理完、标签映射建好之后,下一步是把数据切成训练集、验证集和测试集。很多数据集包在压缩的时候已经把 train/test 分好了,但测试集往往只有文件名没有标签,这种我的建议是:只用它做最终的模型评估,不要碰它。训练过程中需要验证集来判断模型是否过拟合、是否需要调整学习率,于是要从训练集里再切出一块来。
划分比例常见做法是 70% 训练、15% 验证、15% 测试,如果数据集总量比较小(几百张),可以考虑 60% / 20% / 20% 或者直接用 K 折交叉验证。划分的时候有一个关键约束:要按类别分层抽样,保证每一类在三个集合里的比例大致相同。
import os import shutil import random from collections import defaultdict random.seed(42) dataset_root = "fruits_dataset" output_root = "fruits_split" # 按类别收集所有图片路径 class_files = defaultdict(list) for root, dirs, files in os.walk(dataset_root): for fname in files: if not fname.lower().endswith((".jpg", ".jpeg", ".png")): continue label = os.path.basename(root) class_files[label].append(os.path.join(root, fname)) train_ratio, val_ratio, test_ratio = 0.7, 0.15, 0.15 for label, files in class_files.items(): random.shuffle(files) # 先打乱再切分 n_total = len(files) n_train = int(n_total * train_ratio) n_val = int(n_total * val_ratio) train_files = files[:n_train] val_files = files[n_train:n_train + n_val] test_files = files[n_train + n_val:] # 写入目标目录 for split_name, split_files in [("train", train_files), ("val", val_files), ("test", test_files)]: dest_dir = os.path.join(output_root, split_name, label) os.makedirs(dest_dir, exist_ok=True) for src_path in split_files: fname = os.path.basename(src_path) shutil.copy2(src_path, os.path.join(dest_dir, fname)) print("划分完成")random.seed(42)固定随机种子,保证每次运行脚本划分结果完全一致,这在复现实验时非常重要。想象一下你把数据集划分脚本跑了两遍,两次切出来的验证集不一样,模型效果对比就失去了意义。shutil.copy2是复制文件并保留元数据,如果你磁盘空间紧张,可以把copy2换成move,但移动操作有风险——如果中途脚本崩了,原文件被挪走了一半,重新整理会很痛苦。我更推荐先复制、确认无误后再手动删除原始目录。
4.2 图片尺寸、归一化参数与数据增强的选择
划分完目录结构,接下来写 PyTorch 的 Dataset 类和预处理流水线。关键的决策点是:图片要不要统一 resize?用什么尺寸?要不要做数据增强?
水果图片数据集里每张图的尺寸往往不一样,有的 400×400,有的可能是 1200×800。深度学习模型要求输入张量形状一致,所以必须 resize。尺寸选择取决于你用的预训练模型:ResNet 系列常用 224×224,EfficientNet 系列有的用 240×240 或 260×260,Vision Transformer 则常见 224×224 或 384×384。尺寸不是越大越好——更大的尺寸意味着更多的计算量,但信息量并不一定线性增加。我一般先用 224×224 跑通基线,再对比 384×384 看收益,如果提升不到 1 个点,就维持 224 省时间。
归一化参数需要用数据集的均值和标准差。ImageNet 预训练模型的官方归一化参数是mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225],如果你的模型用的是 ImageNet 预训练权重,这个参数直接用,不需要自己算。如果你从零训练,才需要跑一遍代码统计自己数据集的均值和标准差。
from torchvision import datasets, transforms from torch.utils.data import DataLoader train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(degrees=10), 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_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) train_dataset = datasets.ImageFolder(root="fruits_split/train", transform=train_transform) val_dataset = datasets.ImageFolder(root="fruits_split/val", transform=val_transform) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4) print(f"训练集样本数: {len(train_dataset)}, 类别数: {len(train_dataset.classes)}") print(f"验证集样本数: {len(val_dataset)}")ImageFolder会自动读取每个子文件夹的名称作为类别标签,并按字母顺序映射为索引。RandomHorizontalFlip和RandomRotation是轻量级增强,不会改变图片语义;ColorJitter对水果这种颜色是判别性特征的场景要慎用——把苹果的红色调偏了 0.5,模型可能就认不出来了。我的经验是增强强度从小到大逐步加,先在验证集上看效果,而不是一上来就开全套增强。
num_workers是数据加载的并行进程数,Windows 上建议设为 0 或 2,设太大会因为多进程与 CUDA 交互产生奇怪的报错;Linux 上设 4 到 8 都没问题。
4.3 用一个预训练 ResNet 快速跑通基线
数据集和 DataLoader 就绪后,最省力的方案是加载在 ImageNet 上预训练好的 ResNet18,把最后一层全连接换成自己的类别数,然后微调全部参数或只训最后一层。这个方案对水果分类这类任务通常能在一两百个 epoch 内拿到很高的准确率,原因是 ImageNet 本身就包含大量水果类别,预训练模型已经学会了「果皮纹理」「圆形轮廓」「茎叶结构」这类基元特征。
import torch import torch.nn as nn import torch.optim as optim from torchvision import models device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) num_features = model.fc.in_features num_classes = len(train_dataset.classes) model.fc = nn.Linear(num_features, num_classes) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-4) best_val_acc = 0.0 num_epochs = 30 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_dataset) # 验证 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) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() 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_model_fruits.pth") print(f" 保存最佳模型,验证准确率 {val_acc:.4f}") print(f"训练完成,最佳验证准确率: {best_val_acc:.4f}")这个训练脚本里models.ResNet18_Weights.IMAGENET1K_V1是 torchvision 官方推荐的权重加载方式,专门用来替代旧版的pretrained=True参数写法,它会自动下载预训练权重并做相应的预处理。学习率 1e-4 是微调阶段比较安全的起点,如果发现训练集 loss 下降很慢,可以逐步上调到 3e-4 或 5e-4,但不要超过 1e-3,因为预训练权重已经在一个很好的局部最优附近,学习率太大会直接把权重打飞,损失前面几十个 epoch 学到的特征。
5. 避坑清单:从解压到训练的五个翻车现场
5.1 解压后图片全是 0 字节文件
现象:解压过程提示「完成」,但打开目录发现大量 0 KB 的图片文件,Image.verify()直接报OSError: image file is truncated。
原因:下载的 rar 包不完整。网盘工具断点续传或者下载过程中网络波动,导致压缩包损坏,但 7-Zip 可能只报警告而继续解压出部分文件。
解决:回到下载源重新下载,下载完成后先核对文件大小是否与分享页面标注的一致。如果压缩包能解压但某些文件损坏,可以用7z t fruits分类数据集.rar做完整性测试,它会逐个文件校验 CRC 校验码,能精确指出哪些文件损坏。遇到这种情况,宁可重新下载也不要手动修补。
5.2 类别文件夹里混入了desktop.ini和Thumbs.db
现象:文件体检脚本统计非图片文件数发现了几十个desktop.ini或Thumbs.db,这些是 Windows 系统自动生成的隐藏配置文件。
原因:数据集制作者在 Windows 环境下整理文件夹时,系统自动创建了这些文件。打包时没有排除,一起被压缩进来了。
解决:在预处理脚本里做过滤。把non_image_files的处理逻辑从「仅提醒」改成「自动忽略 des_ini、Thumbs.db 这类已知系统文件」:
skip_names = {"desktop.ini", "thumbs.db", ".ds_store", "._*"}._*是 macOS 在 NFS 或 SMB 共享时生成的资源分支文件,同样常见。这些文件对训练无影响,但如果直接喂给ImageFolder会导致类名污染或误读错误,过滤掉就好。
5.3 中文类别名导致训练代码报错
现象:训练脚本报UnicodeDecodeError或者标签显示乱码,尤其是在 Linux 服务器上跑的时候。
原因:压缩包里的中文文件夹名(如「苹果」「香蕉」)在 Windows 下是 GBK 编码,传到 Linux 后变成乱码,Python 默认按 UTF-8 读取时直接失败。
解决:解压时用-mcp=936参数(前面已经提过),或者在建标签映射时统一把中文名映射为英文标识。我个人的习惯是拿到任何中文数据集,第一步就把文件夹名全部改成英文小写,既避免编码问题,也避免后面写代码时频繁切换输入法。批量重命名用一条简单的 shell 命令或者 Python 脚本就能完成:
import os import re def sanitize_dirname(name): name = name.replace("苹果", "apple").replace("香蕉", "banana") name = re.sub(r"[^\w\-_]", "_", name) return name.lower() for d in os.listdir("fruits_dataset"): old_path = os.path.join("fruits_dataset", d) if os.path.isdir(old_path): new_name = sanitize_dirname(d) os.rename(old_path, os.path.join("fruits_dataset", new_name))注意替换规则要提前列全,不要漏掉任何一个类别,否则映射表建立后才发现某类被单独落下了,又得重新跑一遍标签映射。
5.4 训练集和验证集之间出现数据泄露
现象:训练过程验证集准确率停滞在 80% 左右,但测试集准确率却掉到 60% 以下。检查发现验证集和训练集里出现了重复图片。
原因:原始数据集里同一个苹果可能被从不同角度拍了好几张照片,或者数据集中本身就有重复图片(复制粘贴产生)。随机划分时,一张图片的多个副本可能同时被分到训练集和验证集,模型在训练时已经「见过」了验证集的内容,验证准确率虚高。
解决:先对图片做去重,再划分数据集。常见的做法是计算每个文件的 MD5 哈希值或感知哈希,找出内容完全相同或几乎相同的图片。MD5 精确但无法识别缩放、裁剪后内容仍然一样的图;感知哈希(imagehash库)能识别相似图片,但对旋转比较敏感,可以根据实际情况选择:
import hashlib from collections import defaultdict hash_map = defaultdict(list) for root, dirs, files in os.walk("fruits_dataset"): for fname in files: if not fname.lower().endswith((".jpg", ".jpeg", ".png")): continue fpath = os.path.join(root, fname) h = hashlib.md5(open(fpath, "rb").read()).hexdigest() hash_map[h].append(fpath) for h, paths in hash_map.items(): if len(paths) > 1: print(f"重复图片组: {paths}")对每组重复图片,只保留一张,其他移到备份目录。用完再去划分数据集,就不会有泄露问题。这个坑很隐蔽,因为训练 loss 正常下降、验证准确率也很漂亮,直到上线做真实验证才发现模型泛化能力远低于实验指标。
5.5 验证集准确率很高但是新图片识别效果差
现象:训练和验证准确率都在 98% 以上,拿手机拍一张新照片丢给模型,预测结果完全不对。
原因:典型的过拟合到训练集分布。水果数据集里的图片往往都是干净背景、中心构图、光线均匀,而真实场景里水果可能长在树上、被树叶遮挡、光线偏暗、背景杂乱。数据增强强度不足,导致模型学到的特征是「居中且完整的水果」,而不是「水果本身」。
解决:增强验证时的数据多样性。一个稳健的做法是先在验证集上看一眼模型的错误预测,把置信度低于 85% 的样本打印出来,人工确认是模型的失误还是标签本身的错误:
import torch.nn.functional as F model.eval() misclassified = [] with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) probs = F.softmax(outputs, dim=1) conf, preds = torch.max(probs, 1) for i in range(len(labels)): if preds[i] != labels[i]: misclassified.append((conf[i].item(), idx_to_label[preds[i].item()], idx_to_label[labels[i].item()])) for conf, pred_label, true_label in sorted(misclassified, key=lambda x: x[0])[:20]: print(f"置信度 {conf:.3f} | 预测 {pred_label} | 真实 {true_label}")这个脚本帮你定位模型在哪些类别上容易混淆。如果苹果和梨经常互相误判,看特征层面的原因是二者颜色形状都很接近;如果是香蕉被误判为芒果,可能是因为数据集中黄色水果类别的图片背景比较相似。定位到具体混淆对之后,再考虑针对性地增加数据增强——比如对「苹果 vs 梨」的混淆,增加随机旋转的角度范围可能有帮助。
6. 最后的进阶技巧:用混淆矩阵和 Grad-CAM 验证模型到底学到了什么
训练跑完、验证准确率达标,并不代表这个模型可以放心投入使用。两个工具建议每次都做一下:混淆矩阵和 Grad-CAM 热力图。
混淆矩阵能告诉你模型在哪些类别上互相混。准确率是整体指标,但整体指标会掩盖局部的失败——比如模型把 95% 的香蕉都认对了,但 30% 的芒果被认成了橙子,这在实际使用中是完全不同的体验。生成混淆矩阵的代码很短,用 sklearn 自带的工具就行:
from sklearn.metrics import confusion_matrix, classification_report import numpy as np all_preds = [] all_labels = [] model.eval() 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) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm = confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_names=val_dataset.classes))classification_report会输出每个类别的精确率、召回率、F1 分数,比准确率信息量大得多。如果某个类别的召回率明显低,说明模型容易把这一类漏掉——这时候可以考虑增加该类别的样本量或调整分类阈值。
Grad-CAM 热力图则是打开模型的「黑匣子」,让卷积神经网络的注意力区域可视化呈现。这个方法对 ResNet 这类有全局平均池化的模型实现起来特别简单,用torchcam库几行代码就能生成:
from torchcam.methods import GradCAM # 假设模型是 resnet18 cam_extractor = GradCAM(model, target_layer="layer4") model.eval() image, label = val_dataset[0] image = image.unsqueeze(0).to(device) with torch.no_grad(): outputs = model(image) _, predicted = torch.max(outputs, 1) # 提取热力图 activation_map = cam_extractor(predicted.item(), outputs)热力图能直接回答一个问题:模型在判断「苹果」时,是看了苹果的轮廓、颜色还是背景?我曾经遇到过一种情况:模型在西瓜类别上准确率极高,但 Grad-CAM 显示它关注的是图片右下角的桌面条纹,而不是西瓜本身——因为数据集里所有西瓜图片都在同一个木桌上拍摄。这意味着模型学到的不是「西瓜」,而是「木桌上的一团绿」,换到白桌上就废了。热力图能帮你及早发现这类隐藏的偏见。
这两步做完,一个从 rar 压缩包到可验证的模型之间的完整链路就收口了。回想我自己的经历,大多数翻车其实都发生在数据准备阶段——解压乱码、标签映射不一致、验证集泄露,这些问题浪费的时间远超训练本身。看完这篇笔记,希望你能少走这些弯路。也希望帮到你。
本文还有配套的精品资源,点击获取