如何使用torch.utils.data.DataLoader加载Tiny ImageNet-200验证集
Tiny-ImageNet-200验证集加载失败原因
datasets.ImageFolder 要求数据集目录必须按类别子文件夹的结构组织,每个子文件夹内存放对应类别的图片,但官方提供的tiny-imagenet-200验证集不符合这个结构:所有图片直接放在val/images目录下,标签单独存放在val/val_annotations.txt文件中,因此直接调用ImageFolder会报错。
方案1:自定义Dataset类(推荐,无需修改原文件结构)
直接继承PyTorch的Dataset类读取验证集图片和对应标签即可,示例代码如下:
import os from torch.utils.data import Dataset # 如果你用OpenCV读图就替换成对应的读图逻辑,这里示例用PIL from PIL import Image class TinyImageNetValDataset(Dataset): def __init__(self, val_dir, transform=None): self.val_dir = val_dir self.transform = transform self.img_to_label = {} self.class_to_idx = {} # 读取根目录下的wnids.txt生成类别到索引的映射 root_dir = os.path.dirname(val_dir) with open(os.path.join(root_dir, 'wnids.txt'), 'r') as f: for idx, line in enumerate(f.readlines()): class_name = line.strip() self.class_to_idx[class_name] = idx # 读取验证集标注文件生成图片到标签的映射 with open(os.path.join(val_dir, 'val_annotations.txt'), 'r') as f: for line in f.readlines(): parts = line.strip().split('\t') img_name = parts[0] class_name = parts[1] self.img_to_label[img_name] = self.class_to_idx[class_name] self.img_list = list(self.img_to_label.keys()) def __len__(self): return len(self.img_list) def __getitem__(self, idx): img_name = self.img_list[idx] img_path = os.path.join(self.val_dir, 'images', img_name) img = Image.open(img_path).convert('RGB') label = self.img_to_label[img_name] if self.transform: img = self.transform(img) return img, label
调用方式如下,直接替换你原来的ImageFolder逻辑即可:
val_dataset = TinyImageNetValDataset( val_dir=args.val_dir, transform=transforms.Compose([ OpencvResize(256), transforms.CenterCrop(224), ToBGRTensor(), ]) ) # 用DataLoader正常加载 val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)
方案2:修改验证集目录结构适配ImageFolder
如果不想自定义Dataset,可以运行如下脚本把验证集图片移动到对应类别的子文件夹中,适配ImageFolder的目录要求:
import os import shutil val_root = "替换为你的val目录绝对路径" img_dir = os.path.join(val_root, 'images') anno_path = os.path.join(val_root, 'val_annotations.txt') with open(anno_path, 'r') as f: for line in f.readlines(): parts = line.strip().split('\t') img_name = parts[0] class_name = parts[1] # 创建对应类别子文件夹 class_dir = os.path.join(val_root, class_name) os.makedirs(class_dir, exist_ok=True) # 移动图片到对应子文件夹 shutil.move(os.path.join(img_dir, img_name), os.path.join(class_dir, img_name))
脚本运行完成后,你原来的ImageFolder代码就可以正常运行了。
内容的提问来源于stack exchange,提问作者Afsaneh Ebrahimi
相关产品推荐
相关产品推荐

