PyTorch中如何让DataLoader跳过返回空值的无效数据批次
报错原因
PyTorch DataLoader默认的批次处理函数default_collate不支持None类型的拼接,当__getitem__返回(None, None)时,默认拼接逻辑会触发类型错误,导致程序中断。
解决方案
方案1:自定义collate_fn过滤无效样本(推荐,适配任意batch size)
在批次拼接阶段直接过滤无效样本,后续无需修改训练逻辑,且可兼容任意大小的batch size:
- 首先实现自定义拼接函数:
def custom_collate(batch): # 过滤掉所有(None, None)的无效样本 batch = [item for item in batch if item[0] is not None and item[1] is not None] # 过滤后批次为空则返回空张量,也可根据需求调整返回值 if len(batch) == 0: return torch.tensor([]), torch.tensor([]) # 剩余样本走默认拼接逻辑 return torch.utils.data.dataloader.default_collate(batch)
- 初始化DataLoader时传入自定义的
collate_fn:
trainLoader = DataLoader(trainDataset, batch_size=1, shuffle=False, collate_fn=custom_collate)
- 遍历阶段判断空批次直接跳过即可:
for x_batch, y_batch in trainLoader: # 空批次跳过 if x_batch.numel() == 0: continue # 正常执行训练逻辑
方案2:在__getitem__中自动查找有效样本
如果不想修改DataLoader参数,也可以在样本读取阶段自动跳过无效索引,直到找到有效样本:
def __getitem__(self, i): max_idx = len(self) while i < max_idx: x, y = self.getData(i) if x is not None and y is not None: return x, y i += 1 # 遍历到数据集末尾仍无有效样本,可根据需求返回默认值或抛出异常 raise StopIteration
注意该方案需做好索引越界防护,若数据集末尾存在大量无效样本可能触发异常。
方案3:提前预过滤有效索引
如果数据集规模较小,可在Dataset初始化阶段提前遍历所有索引,仅保留可返回有效样本的索引,从根源避免无效样本被读取:
class CustomDataset(Dataset): def __init__(self): # 原有初始化逻辑 self.total_num = # 数据集原始总大小 self.valid_indices = [] # 预遍历所有索引过滤有效项 for idx in range(self.total_num): x, y = self.getData(idx) if x is not None and y is not None: self.valid_indices.append(idx) def __len__(self): # 返回有效样本总数 return len(self.valid_indices) def __getitem__(self, idx): # 映射到真实有效索引 real_idx = self.valid_indices[idx] return self.getData(real_idx)
该方案无需修改DataLoader和遍历逻辑,缺点是初始化阶段需要遍历全量样本,大体积数据集场景下初始化耗时较长。
内容的提问来源于stack exchange,提问作者Martin Perry
相关产品推荐
相关产品推荐

