PyTorch自定义Dataset对象时__len__函数报错求助
问题描述
我正在跟随视频课程进行图像分类任务,编写了如下自定义Dataset类:
from torch.utils.data import Dataset class ChestXRayDataSet(Dataset): def __init__(self, image_dirs, transform): # Initialize the Object def get_image(class_name): # define a function to get images from the provided image directories images = [x for x in os.listdir(image_dirs[class_name]) if x.lower().endswith('png')] print(f'Found {len(images)} Images of Class {class_name}') return images # create a directory to store the images self.images = {} self.class_names = ['normal', 'viral', 'covid'] for c in self.class_names: # store the images in directory with class names self.images[c] = get_image(c) self.image_dirs = image_dirs self.transform = transform def __len__(self): # return the number of images in the dataset num_images = sum([len(self.images[class_name]) for class_name in self.class_names]) return num_images def __getitem__(self, index): class_name = random.choice(self.class_names) index = index % len(self.images[class_name]) # to avoid index out of bound error image_name = self.images[class_name][index] # this is the selected images image_path = os.path.join(self.image_dirs[class_name], image_name) image = Image.open(image_path).convert('RGB') # Finally we return the example, and its index as required by # Dataset class return self.transform(image), self.class_names.index(class_name)
随后通过以下代码创建DataLoader:
batch_size = 6 dl_train = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size)
执行print('Num of Training Batches : ', len(dl_train))时出现如下错误:
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) /tmp/ipykernel_28/1549945295.py in <module> 3 dl_test = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size) 4 ----> 5 print('Num of Training Batches : ', len(dl_train)) 6 #print('Num of Test Batches : ', len(dl_test)) /opt/conda/lib/python3.7/site-packages/torch/utils/data/dataloader.py in __len__(self) 411 self._timeout = loader.timeout 412 self._collate_fn = loader.collate_fn --> 413 self._sampler_iter = iter(self._index_sampler) 414 self._base_seed = torch.empty((), dtype=torch.int64).random_(generator=loader.generator).item() 415 self._persistent_workers = loader.persistent_workers /opt/conda/lib/python3.7/site-packages/torch/utils/data/sampler.py in __len__(self) 240 if self.drop_last: 241 return len(self.sampler) // self.batch_size # type: ignore --> 242 else: 243 return (len(self.sampler) + self.batch_size - 1) // self.batch_size # type: ignore /opt/conda/lib/python3.7/site-packages/torch/utils/data/sampler.py in __len__(self) 67 return iter(range(len(self.data_source))) 68 --> 69 def __len__(self) -> int: 70 return len(self.data_source) 71 TypeError: object of type 'ChestXRayDataSet' has no len()
使用环境:Kaggle Notebook,PyTorch版本1.11.0+cpu。
解决方案
错误根源是**__len__和__getitem__方法被定义在了__init__方法内部**,导致它们不是类的成员方法,PyTorch无法识别到这些方法,进而报错“ChestXRayDataSet has no len()”。
修复步骤:
- 将
__len__和__getitem__方法的缩进调整,使其成为类的直接成员,而不是嵌套在__init__里。 - 补充导入缺失的模块:代码中使用了
random、os和Image,需要在开头添加对应导入语句。
修复后的完整代码:
import os import random from PIL import Image from torch.utils.data import Dataset class ChestXRayDataSet(Dataset): def __init__(self, image_dirs, transform): # Initialize the Object def get_image(class_name): # define a function to get images from the provided image directories images = [x for x in os.listdir(image_dirs[class_name]) if x.lower().endswith('png')] print(f'Found {len(images)} Images of Class {class_name}') return images # create a directory to store the images self.images = {} self.class_names = ['normal', 'viral', 'covid'] for c in self.class_names: # store the images in directory with class names self.images[c] = get_image(c) self.image_dirs = image_dirs self.transform = transform def __len__(self): # return the number of images in the dataset num_images = sum([len(self.images[class_name]) for class_name in self.class_names]) return num_images def __getitem__(self, index): class_name = random.choice(self.class_names) index = index % len(self.images[class_name]) # to avoid index out of bound error image_name = self.images[class_name][index] # this is the selected images image_path = os.path.join(self.image_dirs[class_name], image_name) image = Image.open(image_path).convert('RGB') # Finally we return the example, and its index as required by # Dataset class return self.transform(image), self.class_names.index(class_name)
补充说明:
- PyTorch的
Dataset抽象类要求必须实现__len__和__getitem__两个类方法,否则无法被DataLoader正确使用。 - 当前
__getitem__的实现存在逻辑问题:每次调用都会随机选择类别,这会导致数据加载顺序完全随机,可能出现样本重复加载或从未加载的情况。如果需要遍历所有样本,建议预先构建包含所有样本路径和标签的列表,再根据index直接取对应样本。
内容的提问来源于stack exchange,提问作者new2Ubuntu
相关产品推荐
相关产品推荐

