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

如何保存/加载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]

额外注意事项

  1. 测试时要保证输入图片的预处理和训练完全一致,比如归一化、尺寸等,若训练时用了其他Transform,测试时也要同步添加
  2. 测试代码中有一处小错误:model.eval()应改为loadedModel.eval(),虽不影响结果,但建议修正

内容的提问来源于stack exchange,提问作者Niaisenif

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.01 14:29:52