PyTorch Ignite模型Checkpoint保存加载后单张图片推理报错的解决方法
解决PyTorch Ignite保存模型后单图推理的维度不匹配错误
问题背景
我最近用PyTorch Ignite的Checkpoint系统保存自定义CNN模型,批量测试数据能正常运行,但单张图片推理时触发了这个错误:
RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x841 and 26912x3364)
错误原因
这个问题不是模型保存/加载的问题,而是输入维度不匹配导致的:
- 训练时输入的是批量数据,shape为
(batch_size, 3, 256, 256),经过卷积池化后得到(batch_size, 32, 29, 29),nn.Flatten()会把它展平成(batch_size, 32*29*29),刚好匹配线性层的输入维度32*29*29。 - 但单张图片输入时,传入的是
(3, 256, 256)(没有batch维度),经过卷积池化后变成(32, 29, 29),nn.Flatten()默认从第1维开始展平,结果是(32, 29*29) = (32, 841),这和线性层期望的26912(322929)输入维度完全不匹配,所以报错。
解决方案
最简单的办法是让单图输入符合模型训练时的格式,下面提供两种可行方式:
方式1:推理前手动添加batch维度
在传入模型前,用unsqueeze(0)给测试张量增加第0维(batch维度):
testTensor = testTransforms(testTensor) # 增加batch维度,shape从(3,256,256)变为(1,3,256,256) testTensor = testTensor.unsqueeze(0) print(testTensor.shape) # 现在是torch.Size([1, 3, 256, 256]) with torch.no_grad(): output = loadedModel(testTensor) # 如果需要去掉batch维度,可以用output.squeeze(0) print(output)
方式2:修改模型forward方法自动兼容
如果不想每次手动加维度,可以修改模型的forward方法,自动检测并处理无batch维度的输入:
def forward(self, x): # 检测输入是否为3D(无batch维度:C,H,W) if len(x.shape) == 3: x = x.unsqueeze(0) # 增加batch维度变为(1,C,H,W) x = self.poolingStack(x) x = self.flatten(x) logits = self.linear_relu_stack(x) # 如果原始输入是3D,返回时去掉batch维度 if len(logits.shape) == 2 and logits.shape[0] == 1: logits = logits.squeeze(0) return logits
验证修改
修改后,单张图片输入就能正常推理了,同时批量输入也不会受影响——因为批量输入的shape是4D,不会触发len(x.shape)==3的分支。
额外提醒
PyTorch的模型默认都是为批量输入设计的,输入张量的第一维是batch size,日常推理时尽量保持这个规范,能避免很多维度相关的错误。
内容的提问来源于stack exchange,提问作者Niaisenif
相关产品推荐
相关产品推荐

