使用torchvision加载Celeba数据集时遭遇加载错误求助
解决CelebA数据集加载的MD5校验与文件找不到问题
一、核对手动下载后的文件结构
torchvision的CelebA数据集要求严格的目录结构,必须匹配以下格式:
./ └── celeba/ ├── img_align_celeba/ │ ├── 000001.jpg │ ├── 000002.jpg │ └── ... (所有图片文件) ├── list_attr_celeba.txt ├── list_bbox_celeba.txt ├── list_eval_partition.txt └── list_landmarks_align_celeba.txt
对照你的本地文件:
- 图片文件夹必须命名为
img_align_celeba,不能用其他名称 - 所有txt文件要直接放在
celeba根目录下,不能嵌套在子文件夹中
二、修复自动下载的MD5校验错误
如果download=True时MD5报错,是因为官方源文件更新或镜像同步问题,可按以下步骤处理:
- 获取本地文件的实际MD5值:
- Linux/macOS终端执行:
md5sum img_align_celeba.zip list_attr_celeba.txt - Windows PowerShell执行:
Get-FileHash img_align_celeba.zip -Algorithm MD5
- Linux/macOS终端执行:
- 修改torchvision源码中的MD5值:
找到你Python环境中torchvision/datasets/celeba.py文件,定位到MD5字典,替换为你本地文件的实际MD5值,再重新运行download=True
三、自定义Dataset绕开torchvision校验(终极方案)
如果上述方法无效,直接自定义数据集加载逻辑,完全绕开torchvision的校验机制:
import os import pandas as pd from PIL import Image from torch.utils.data import Dataset class CustomCelebA(Dataset): def __init__(self, root, split='train', target_type='attr'): self.root = root self.img_dir = os.path.join(root, 'celeba', 'img_align_celeba') self.attr_path = os.path.join(root, 'celeba', 'list_attr_celeba.txt') self.split_path = os.path.join(root, 'celeba', 'list_eval_partition.txt') # 加载数据集划分信息 split_df = pd.read_csv(self.split_path, sep='\s+', header=None, names=['img_id', 'split']) split_map = {'train':0, 'valid':1, 'test':2} self.img_ids = split_df[split_df['split'] == split_map[split]]['img_id'].tolist() # 加载属性标签 attr_df = pd.read_csv(self.attr_path, sep='\s+', skiprows=1) self.attr_names = attr_df.columns[1:] self.attrs = attr_df.loc[self.img_ids, self.attr_names].values # 将原始标签的-1转为0(匹配torchvision的默认处理逻辑) self.attrs = (self.attrs + 1) // 2 def __len__(self): return len(self.img_ids) def __getitem__(self, idx): img_id = self.img_ids[idx] img_path = os.path.join(self.img_dir, img_id) img = Image.open(img_path).convert('RGB') attr = self.attrs[idx] return img, attr # 实例化自定义数据集 celeba_train_dataset = CustomCelebA('./', split='train') celeba_valid_dataset = CustomCelebA('./', split='valid') celeba_test_dataset = CustomCelebA('./', split='test')
四、常见排查点
- 检查图片数量:训练集162770张、验证集19867张、测试集19962张,确保解压后数量匹配
- 验证txt文件格式:用文本编辑器打开,确认没有乱码、行尾符错误或内容缺失
- Windows用户注意路径大小写:torchvision对路径大小写敏感,确保文件夹和文件名与要求完全一致
内容的提问来源于stack exchange,提问作者Sahl Ebrahim
相关产品推荐
相关产品推荐

