DataLoader加载含input_ids列的数据集时因张量尺寸不一致触发RuntimeError
DataLoader加载含input_ids列的数据集时因张量尺寸不一致触发RuntimeError
问题场景
我最近碰到了一个棘手的问题:手里有个包含input_ids列的数据集,用PyTorch的DataLoader加载时,只要batch_size设为2及以上,遍历过程中就会触发RuntimeError,但把batch_size改成1就完全正常。
先贴一下我的核心代码:
train_batch_size = 2 eval_dataloader = DataLoader(val_dataset, batch_size=train_batch_size) # 打印dataloader的批次数量 print(len(eval_dataloader)) >>> 1623
当我尝试遍历这个dataloader时:
for step, batch in enumerate(eval_dataloader): print(step) >>> 1,2... ,1621 # 还没遍历完所有批次就报错中断了
报错栈最终指向了torch.stack操作,明确提示张量尺寸不匹配:
RuntimeError: stack expects each tensor to be equal size, but got [212] at entry 0 and [1024] at entry 1
我一开始以为是最后一个批次的样本数量不足导致的,试着加上drop_last=True参数,但报错依然存在,就连训练集的dataloader也出现了一模一样的问题。
问题根源
折腾了半天终于搞清楚:数据集里混进了长度不是1024的序列!当batch_size大于1时,DataLoader默认的default_collate函数会尝试把批次内所有input_ids张量堆叠成一个大张量,但如果其中有样本的序列长度和其他样本不一致,堆叠操作就会直接失败;而batch_size=1时,不需要堆叠多个张量,自然就不会触发这个错误。
解决步骤
按照下面的方法一步步操作就能搞定:
- 定位异常样本:把batch_size设为1、关闭shuffle,遍历数据集打印每个样本的
input_ids形状,找出长度不符合要求的样本:temp_dataloader = DataLoader(val_dataset, batch_size=1, shuffle=False) for idx, batch in enumerate(temp_dataloader): print(idx, batch['input_ids'].shape) - 过滤数据集:移除那些长度不是1024的异常样本,比如用列表推导式筛选:
filtered_dataset = [item for item in val_dataset if len(item['input_ids']) == 1024] - 重新初始化DataLoader:用过滤后的数据集创建新的dataloader,再遍历就不会报错了:
eval_dataloader = DataLoader(filtered_dataset, batch_size=train_batch_size, drop_last=True)
备注:内容来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

