PyTorch在保存模型时会保存批量大小:我该如何保存模型,才能一次输入一张图像?

人工智能 2026-07-10

我在使用PyTorch Ignite以便更容易地获取指标,因此我在使用他们的检查点系统来保存我的模型。
我想在用户界面中实现一次测试一张图片,但当我尝试这样做时,我会得到一个

RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x841 and 26912x3364)w

当我输入一个完整的批次时。
如何保存/加载模型,以便能够一次输入一张图片?

模型代码:

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(64 * 30 * 30, 64 * 30 * 30),
            # nn.ReLU(),
            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



# hyperparameters
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()


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))

最小可复现示例(Colab文件)

# -*- coding: utf-8 -*-
"""Untitled1.ipynb

Automatically generated by Colab.

Original file is located at
    https://colab.research.google.com/drive/1E4o0XN8npMv8W6Tuj372ZI9vfpl5s88K
"""

!pip install pytorch-ignite
import kagglehub

import os
import torch
import torch.nn as nn
import torchvision.models as models
from torch.utils.data import DataLoader, Dataset
import torchvision.transforms as Transforms
from torchvision.io import decode_image
import pandas as pd

from ignite.engine import *
from ignite.handlers import *
from ignite.metrics import *
from ignite.metrics.clustering import *
from ignite.metrics.regression import *
from ignite.utils import *
from ignite.contrib.handlers import TensorboardLogger, global_step_from_engine


path = kagglehub.dataset_download("doctorstrange420/real-and-fake-ai-generated-art-images-dataset")

# taken from colab so path might be wrong
path = "/root/.cache/kagglehub/datasets/doctorstrange420/real-and-fake-ai-generated-art-images-dataset/versions/1/Data"


# ==================================================================================================

realIMG = os.listdir(path + "/REAL")
fakeIMG = os.listdir(path + "/FAKE")

imgArr = realIMG.copy()
imgArr = imgArr + fakeIMG

labelArr = [0 for x in range(0, len(realIMG))]
labelArr = labelArr + [1 for x in range(0, len(fakeIMG))]


dataDict = {"img":imgArr, "label": labelArr}
train_df = pd.DataFrame.from_dict(dataDict, orient="columns")
train_df.info()

train_df.to_csv(r'training_set.csv', index = False, header=True)

# ==================================================================================================

class ImgDataset(Dataset):
    def __init__(self, annotations_file, img_dir, transform=None, target_transform=None):
        self.img_labels = pd.read_csv(annotations_file)
        self.img_dir = img_dir
        self.transform = transform
        self.target_transform = target_transform

    def __len__(self):
        return len(self.img_labels)

    def __getitem__(self, idx):
        label = self.img_labels.iloc[idx, 1]
        if label == 0:
          dir = "REAL"
        else :
          dir = "FAKE"
        img_path = os.path.join(self.img_dir, dir, self.img_labels.iloc[idx, 0])
        image = decode_image(img_path).float().to(device)
        if self.transform:
            image = self.transform(image)
        return image, label

# ==================================================================================================


class myCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.poolingStack = nn.Sequential(
            nn.Conv2d(3, 8, kernel_size=9), # => 8*(256-(9/2))² == 8*248*248
            nn.ReLU(),
            nn.MaxPool2d(2), # => 8 * 124 * 124
            nn.Conv2d(8, 16, kernel_size=5), # => 16*120*120
            nn.ReLU(),
            nn.MaxPool2d(2), # => 16 * 60 * 60
            nn.Conv2d(16, 32, kernel_size=3), # => 32 * 58 * 58
            nn.ReLU(),
            nn.MaxPool2d(2) # => 32 * 29 * 29
        )
        self.flatten = nn.Flatten()
        self.linear_relu_stack = nn.Sequential(
            # nn.Linear(64 * 30 * 30, 64 * 30 * 30),
            # nn.ReLU(),
            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


# ==================================================================================================


# hyperparameters
num_epochs = 15
learning_rate = 1e-5
batch_size = 64
shuffle = True

transforms = None

dataset = ImgDataset("training_set.csv", path, transforms)
train_set, validation_set = torch.utils.data.random_split(dataset,[17500,4142])
train_dataloader = DataLoader(dataset=train_set, shuffle=shuffle, batch_size=batch_size)
test_dataloader = DataLoader(dataset=validation_set, shuffle=shuffle, batch_size=batch_size)


model = myCNN().to(device)


loss_func = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)

# ==================================================================================================


trainer = create_supervised_trainer(model, optimizer, loss_func, device)


def binary_one_hot_output_transform(output):
    # print(output)
    y_pred, y = output
    y_pred = torch.sigmoid(y_pred).round().long()
    # y_pred = utils.to_onehot(y_pred, 2)
    y = y.long()

    # print(y_pred, y)
    return y_pred, y

val_metrics = {
    "accuracy": Accuracy()
}

train_evaluator = create_supervised_evaluator(model, metrics=val_metrics, device=device)
val_evaluator = create_supervised_evaluator(model, metrics=val_metrics, device=device)


# How many batches to wait before logging training status
log_interval = 100

@trainer.on(Events.ITERATION_COMPLETED(every=log_interval))
def log_training_loss(engine):
    print(f"Epoch[{engine.state.epoch}], Iter[{engine.state.iteration}] Loss: {engine.state.output:.2f}")


@trainer.on(Events.EPOCH_COMPLETED)
def log_training_results(trainer):
    train_evaluator.run(train_dataloader)
    metrics = train_evaluator.state.metrics
    print(f"Training Results - Epoch[{trainer.state.epoch}]\n\tAvg accuracy: {metrics['accuracy']:.2f}\n\tAvg loss: {metrics['loss']:.2f}\n\tAvg recall: {metrics['recall']}\n\tAvg precision: {metrics['precision']}\n\tAvg f1: {metrics['f1']:.2f}")
    print(f"\tcm : {metrics['cm']}")


@trainer.on(Events.EPOCH_COMPLETED)
def log_validation_results(trainer):
    val_evaluator.run(test_dataloader)
    metrics = val_evaluator.state.metrics
    print(f"Training Results - Epoch[{trainer.state.epoch}]\n\tAvg accuracy: {metrics['accuracy']:.2f}\n\tAvg loss: {metrics['loss']:.2f}\n\tAvg recall: {metrics['recall']}\n\tAvg precision: {metrics['precision']}\n\tAvg f1: {metrics['f1']:.2f}")
    print(f"\tcm : {metrics['cm']}")

# Score function to return current value of any metric we defined above in val_metrics
def score_function(engine):
    return engine.state.metrics["accuracy"]



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})

tb_logger = TensorboardLogger(log_dir="tb-logger")

tb_logger.attach_output_handler(
    trainer,
    event_name=Events.ITERATION_COMPLETED(every=log_interval),
    tag="training",
    output_transform=lambda loss: {"batch_loss": loss},
)

for tag, evaluator in [("training", train_evaluator), ("validation", val_evaluator)]:
    tb_logger.attach_output_handler(
        evaluator,
        event_name=Events.EPOCH_COMPLETED,
        tag=tag,
        metric_names="all",
        global_step_transform=global_step_from_engine(trainer),
    )

trainer.run(train_dataloader, max_epochs=num_epochs)
tb_logger.close()

第二部分:

loadedModel = myCNN().to(device)
loadedModel.load_state_dict(torch.load("/content/checkpoint/best_model_14_accuracy=0.8621.pt", map_location=device)) # change based on saved checkpoint
loadedModel.eval()

testImage = "/content/Mario and Luigi eat at Burger King - 0-0-05.jpeg" # put whatever

testTensor = decode_image(testImage).float().to(device)

testTransforms = Transforms.Compose([
        Transforms.Resize((256, 256))
    ])

testTensor = testTransforms(testTensor)

with torch.no_grad():
  model.eval()
  print(loadedModel(testTensor))

解决方案

根据 decode_image 的文档,decode_image(testImage) 的返回值是一个形状为 (image_channels, image_height, image_width) 的三维张量,但你的模型输入期望的是形状为 (batch_size, image_channels, image_height, image_width) 的四维张量。

你可以很容易地通过在张量上再添加一个维度来修复这个问题,使用 .unsqueeze(0),它将创建一个形状为 (1, image_channels, image_height, image_width) 的张量,其中 1 是批量大小:

testTensor = decode_image(testImage).float().to(device).unsqueeze(0)  # <- here
testTransforms = Transforms.Compose([Transforms.Resize((256, 256))])
testTensor = testTransforms(testTensor)
with torch.no_grad():
    print(loadedModel(testTensor))
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。

相关文章