使用PyTorch自定义Dataset加载图像数据时DataLoader运行报错如何解决?
报错原因分析
- 核心问题:自定义的
data_test数据集类没有实现PyTorchDataset要求的必填__len__方法。DataLoader运行时需要通过该方法获取数据集总样本量,完成批次划分、遍历终止判断等逻辑,缺失该方法会直接触发运行异常。 - 次要排查点:检查代码缩进是否规范,
__init__、__getitem__方法必须缩进在data_test类的作用域下,否则会被识别为全局独立函数,Dataset同样会判定核心方法缺失。 - 潜在隐藏问题:
cv2.imread默认返回BGR通道顺序的numpy数组,若你传入的transforms是PyTorch官方针对PIL图像设计的变换组合,可能会因为输入格式不匹配触发后续额外报错。
解决方法
首先补全数据集类的必填方法,适配加载逻辑,修复后代码如下:
import cv2 import glob from torch.utils.data import Dataset, DataLoader class data_test(Dataset): def __init__(self, data_root, transform=None): self.data_image = glob.glob(data_root + '/*.jpg') self.transform = transform # 新增必填__len__方法,返回数据集总样本数 def __len__(self): return len(self.data_image) def __getitem__(self, index): data_image_path = self.data_image[index] image_data = cv2.imread(data_image_path, -1) # 可选:BGR转RGB,适配PyTorch默认的图像处理逻辑 if len(image_data.shape) == 3 and image_data.shape[2] == 3: image_data = cv2.cvtColor(image_data, cv2.COLOR_BGR2RGB) if self.transform: image_data = self.transform(image_data) return image_data
其他验证事项:
- 确认传入的
train_dataset路径真实存在,且路径下至少有1张后缀为.jpg的图片,避免glob返回空列表导致加载异常。 - 若使用的
transforms包含PIL专属变换,可将读取后的numpy数组转为PIL Image(Image.fromarray(image_data)),或直接使用torchvision.io.read_image读取图片适配张量变换逻辑。
内容的提问来源于stack exchange,提问作者Feona
相关产品推荐
相关产品推荐

