PyTorch张量错误:输入与目标batch_size不匹配问题求助
石头剪刀布CNN训练错误排查与解决
核心问题
你碰到的ValueError: Expected input batch_size (1) to match target batch_size (2),结合打印出的target尺寸为torch.Size([1,0]),本质是标签张量为空——目标数据(outputs)没有有效内容,导致框架计算时出现batch维度匹配异常。
排查修复步骤
- 检查
train_batch_creator(注意拼写应为train_batch_creator)的标签生成逻辑:- 确认每个图像对应的标签(0=paper、1=rock、2=scissors)是否正确赋值,有没有漏写、误赋值为空的情况。比如生成paper标签时是不是把0写成了空值,或者标签列表被意外清空。
- 打印
shuffled_outputs的具体内容,查看是否存在大量空元素,或者和shuffled_inputs长度不匹配。正常情况下两者长度必须完全一致。
- 验证洗牌代码的执行结果:
- 在洗牌后加一段代码,检查输入和标签列表的长度:
如果长度不一致,说明洗牌前的print(len(shuffled_inputs), len(shuffled_outputs))inputs和outputs长度就不匹配,得回到数据生成环节找问题。
- 在洗牌后加一段代码,检查输入和标签列表的长度:
- 检查
img_configs函数的图像处理逻辑:- 确认函数没有只过滤掉部分图像,却没同步删除对应标签,导致标签列表和输入列表长度不一致。
- 修复标签张量维度:
- 即便标签生成正确,也要确保标签维度是
[batch_size](比如输入batch是1时,标签应为torch.Size([1])),而非[1,0]这种异常维度。训练前可以用torch.squeeze()去除空维度,或者重新整理标签形状。
- 即便标签生成正确,也要确保标签维度是
补充说明
你用的洗牌代码逻辑没问题,不用修改。重点要盯紧数据生成环节的输入与标签配对一致性,确保每个输入图像都对应有效标签。
内容的提问来源于stack exchange,提问作者Edin Tresnjic
相关产品推荐
相关产品推荐

