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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 06:40:25