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
相关产品推荐
相关产品推荐

