PyTorch Ignite模型适配单张图片输入的保存与加载方案:解决形状不匹配RuntimeError
Fix RuntimeError for Single Image Inference with PyTorch Ignite Saved Model
你遇到的RuntimeError: mat1 and mat2 shapes cannot be multiplied问题,核心原因是单张图片输入缺少batch维度。PyTorch的卷积层和全连接层默认期望输入是4维张量([batch_size, channels, height, width]),但你传入的测试张量是3维的([3,256,256]),这导致模型对维度的理解错位,最终在全连接层触发形状不匹配的错误。
以下是具体的解决方案:
解决方案1:手动给测试图片添加batch维度
这是最直接的修复方式,只需要在输入模型前给测试张量添加一个batch维度,让输入形状和训练时的预期一致:
loadedModel = myCNN().to(device) loadedModel.load_state_dict(torch.load("/content/checkpoint/best_model_14_accuracy=0.8621.pt", map_location=device)) loadedModel.eval() testTensor = decode_image(testImage).float().to(device) testTransforms = Transforms.Compose([ Transforms.Resize((256, 256)) ]) testTensor = testTransforms(testTensor) # 关键:添加batch维度,形状变为[1, 3, 256, 256] testTensor = testTensor.unsqueeze(0) print(testTensor.shape) with torch.no_grad(): output = loadedModel(testTensor) print(output) # 可选:如果需要去掉batch维度,用squeeze(0) # output = output.squeeze(0)
解决方案2:修改模型Forward方法,自动适配输入维度
如果你希望模型能同时兼容3维(单张无batch)和4维(批量)输入,可以在forward方法中自动处理维度:
class myCNN(nn.Module): def __init__(self): super().__init__() self.poolingStack = nn.Sequential( nn.Conv2d(3, 8, kernel_size=9), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(8, 16, kernel_size=5), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(16, 32, kernel_size=3), nn.ReLU(), nn.MaxPool2d(2) ) self.flatten = nn.Flatten() self.linear_relu_stack = nn.Sequential( nn.Linear(32 * 29 * 29, 4* 29 * 29), nn.ReLU(), nn.Dropout(0.2), nn.Linear(4 * 29 * 29, 29 * 29), nn.ReLU(), nn.Dropout(0.2), nn.Linear(29*29, 2) ) def forward(self, x): # 检查输入是否为3维,自动添加batch维度 if len(x.shape) == 3: x = x.unsqueeze(0) x = self.poolingStack(x) x = self.flatten(x) logits = self.linear_relu_stack(x) # 若原始输入是3维,输出自动去掉batch维度 if len(logits.shape) == 2 and logits.shape[0] == 1: logits = logits.squeeze(0) return logits
修改后,无论是传入3维还是4维张量,模型都能正确处理,无需手动调整维度。
模型保存/加载的额外注意事项
你的PyTorch Ignite保存和加载流程是正确的,再补充几个关键点:
- 加载模型时用
map_location=device确保张量和当前设备一致,避免跨设备错误 - 加载后必须调用
loadedModel.eval(),切换到评估模式,关闭Dropout等训练专属层 - 测试时的预处理逻辑(比如Resize尺寸)必须和训练集完全一致,否则会导致特征不匹配
内容的提问来源于stack exchange,提问作者Niaisenif
相关产品推荐
相关产品推荐

