PyTorch中__getitem__返回None给DataLoader会发生什么?
PyTorch中Dataset返回None给DataLoader的结果
你的猜测不对,实际情况是直接触发错误,训练无法继续,具体细节如下:
- PyTorch默认的
default_collate函数会把批次内的元素打包成张量,而None无法被转换为张量,也没法和其他正常样本拼接,所以会直接抛出TypeError。 - 举个实际代码例子验证:
from torch.utils.data import Dataset, DataLoader import torch class BadDataset(Dataset): def __len__(self): return 5 def __getitem__(self, idx): if idx == 2: return None return idx # 使用默认collate函数的DataLoader loader = DataLoader(BadDataset(), batch_size=2) for batch in loader: print(batch)
运行后会在处理到第3个样本时报错:TypeError: expected sequence of integers or floats, got NoneType。
如果想要实现「跳过无效样本、批次大小自动调整」的效果,不能让__getitem__返回None,得换两种方式处理:
- 方式一:在Dataset内部过滤无效样本,比如遇到无效数据时递归获取下一个有效样本(注意要加边界判断避免死循环)。
- 方式二:自定义
collate_fn,在函数里过滤掉None的样本,再返回剩余元素组成的批次(此时批次大小可能小于设定值),示例代码:
def custom_collate(batch): # 过滤批次里的None样本 batch = [item for item in batch if item is not None] if not batch: return None # 空批次可根据业务需求返回空张量或跳过 return torch.tensor(batch) # 使用自定义collate函数的DataLoader loader = DataLoader(BadDataset(), batch_size=2, collate_fn=custom_collate) for batch in loader: if batch is not None: print(batch)
这种情况下训练可以继续,无效样本会被跳过,批次大小可能会缩小。
内容的提问来源于stack exchange,提问作者Ink
相关产品推荐
相关产品推荐

