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

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时,不需要堆叠多个张量,自然就不会触发这个错误。

解决步骤

按照下面的方法一步步操作就能搞定:

  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)
    
  2. 过滤数据集:移除那些长度不是1024的异常样本,比如用列表推导式筛选:
    filtered_dataset = [item for item in val_dataset if len(item['input_ids']) == 1024]
    
  3. 重新初始化DataLoader:用过滤后的数据集创建新的dataloader,再遍历就不会报错了:
    eval_dataloader = DataLoader(filtered_dataset, batch_size=train_batch_size, drop_last=True)
    

备注:内容来源于stack exchange,提问作者Penguin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.22 07:04:30