如何保存/加载PyTorch模型以支持单张图片输入推理?
问题描述
我用PyTorch Ignite的checkpoint系统保存模型,训练时输入完整批次数据运行正常,但单张图片测试时报错:
RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x841 and 26912x3364)
想调整保存/加载逻辑,让模型支持单张图片输入。
相关代码
模型定义
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): x = self.poolingStack(x) x = self.flatten(x) logits = self.linear_relu_stack(x) return logits # 超参数 num_epochs = 15 learning_rate = 1e-5 batch_size = 64 shuffle = True transforms = None model = myCNN().to(device)
模型保存代码
@trainer.on(Events.EPOCH_COMPLETED) def log_training_results(trainer): train_evaluator.run(train_dataloader) @trainer.on(Events.EPOCH_COMPLETED) def log_validation_results(trainer): val_evaluator.run(test_dataloader) model_checkpoint = ModelCheckpoint( "checkpoint", n_saved=2, filename_prefix="best", score_function=score_function, score_name="accuracy", global_step_transform=global_step_from_engine(trainer), ) val_evaluator.add_event_handler(Events.COMPLETED, model_checkpoint, {"model": model})
模型加载及测试代码
loadedModel = myCNN().to(device) loadedModel.load_state_dict(torch.load("/content/checkpoint/best_model_14_accuracy=0.8621.pt", map_location=device)) loadedModel.eval() testImage = "/content/Mario and Luigi eat at Burger King - 0-0-05.jpeg" testTensor = decode_image(testImage).float().to(device) testTransforms = Transforms.Compose([ Transforms.Resize((256, 256)) ]) testTensor = testTransforms(testTensor) print(testTensor.shape) # => torch.Size([3, 256, 256]) with torch.no_grad(): print(loadedModel(testTensor))
解决方案
这个问题不是模型保存/加载的问题,而是输入数据维度不匹配导致的:
- 训练时DataLoader输出批次数据,维度为
(batch_size, channels, height, width)(比如(64,3,256,256)) - 单张测试时输入维度是
(3,256,256),缺少batch维度,导致卷积后特征展平后,全连接层输入维度不匹配
解决方法只需给单张图片Tensor添加batch维度,有两种方式:
方法1:使用unsqueeze添加维度
修改测试代码的输入部分:
# 原测试Tensor形状是(3,256,256) testTensor = testTransforms(testTensor) # 添加batch维度,变成(1,3,256,256) testTensor = testTensor.unsqueeze(0) with torch.no_grad(): print(loadedModel(testTensor))
方法2:使用[None]快速添加维度
这是PyTorch语法糖,效果和unsqueeze(0)一致:
testTensor = testTransforms(testTensor)[None]
额外注意事项
- 测试时要保证输入图片的预处理和训练完全一致,比如归一化、尺寸等,若训练时用了其他Transform,测试时也要同步添加
- 测试代码中有一处小错误:
model.eval()应改为loadedModel.eval(),虽不影响结果,但建议修正
内容的提问来源于stack exchange,提问作者Niaisenif
相关产品推荐
相关产品推荐

