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

PyTorch自定义数据集仅加载满足条件数据时遇TypeError问题求助

解决PyTorch Dataset中返回None导致的TypeError问题

你遇到的这个错误核心原因很明确:你的__getitem__方法在某些索引下没有返回任何有效数据,直接返回了None,而PyTorch的default_collate函数无法处理NoneType对象。

看一下你的代码逻辑:你想要加载第三列为'y'的数据,但写的判断条件是if (self.annotations.iloc[index, 2] == 'n')——这完全写反了!而且当条件不满足时(也就是第三列是'y'的时候),方法没有任何return语句,自然就返回了None,Dataloader在整理batch的时候遇到这些None就会抛出错误。

另外还有个问题:即使你把条件改对,__len__返回的是原始CSV的总行数,而Dataloader会遍历所有索引,这时候如果某些索引对应的数据不符合条件,还是会出现返回None的情况。

正确的解决方案:提前过滤符合条件的数据

最稳妥的方式是在__init__阶段就把第三列为'y'的数据筛选出来,这样后续的__len__和__getitem__都只处理有效数据,避免出现None的情况。修改后的代码如下:

import pandas as pd
import cv2
from torch.utils.data import Dataset

class InterDataset(Dataset):
    def __init__(self, csv_file, mode, root_dir = None, transform = None, run = None):
        # 读取CSV后直接过滤第三列为'y'的数据
        self.annotations = pd.read_csv(csv_file, header = None)
        # 筛选条件:第三列等于'y'
        self.annotations = self.annotations[self.annotations.iloc[:, 2] == 'y'].reset_index(drop=True)
        self.root_dir = root_dir
        self.transform = transform
        self.mode = mode
        self.run = run
        
    def __len__(self):
        # 返回过滤后的有效数据行数
        return len(self.annotations)
    
    def __getitem__(self, index):
        # 现在所有索引对应的都是符合条件的数据,不需要再判断
        if self.mode == 'train':
            img_path = self.annotations.iloc[index, 0]
            image = cv2.imread(img_path, 1)
            # 注意:PyTorch默认图像格式是RGB,而cv2读出来是BGR,建议转一下
            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
            y_label = self.annotations.iloc[index, 1]
            
            if self.transform:
                image = self.transform(image)
                
            if (index+1) % 300 == 0:
                print(f'Loop {index} done')
                
            return image, y_label  # 这里返回元组比列表更规范

为什么这样修改?

  • 提前过滤:在初始化阶段就把无效数据排除,保证__getitem__处理的每一个索引都是有效数据,不会返回None。
  • 重置索引:使用reset_index(drop=True)重置索引,避免原CSV的索引混乱,确保__getitem__的索引和过滤后的DataFrame对应。
  • 图像格式修正:补充了BGR转RGB的步骤,因为PyTorch的大部分图像预处理transform都是基于RGB格式设计的,避免后续出现颜色异常的问题。
  • 返回元组:返回元组比列表更符合PyTorch Dataset的常规做法,Dataloader处理起来也更顺畅。

如果你不想提前过滤(不推荐)

如果出于某些原因必须在__getitem__中动态判断,那你需要确保永远不返回None,可以通过循环找到下一个符合条件的索引,但这种方式效率很低,而且容易导致batch大小不一致,不建议使用。比如:

def __getitem__(self, index):
    while True:
        if self.annotations.iloc[index, 2] == 'y':
            # 处理数据并返回
            img_path = self.annotations.iloc[index,0]
            image = cv2.imread(img_path,1)
            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
            y_label = self.annotations.iloc[index,1]
            if self.transform:
                image = self.transform(image)
            return image, y_label
        # 如果当前索引不符合条件,就取下一个索引
        index = (index + 1) % len(self.annotations)

但这种方式会打乱数据的顺序,而且在数据量很大时会有性能问题,所以还是优先推荐提前过滤的方案。

内容的提问来源于stack exchange,提问作者BlodyRevenge199

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 19:12:38