You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 14:10:35