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

PyTorch迁移学习图像分类模型始终返回第0类索引问题

迁移学习图像分类预测错误问题

使用迁移学习做蚂蚁和蜜蜂的图像分类,直接复制PyTorch官方迁移学习教程的完整代码,在PyCharm中训练并保存模型。加载模型后输入单张蜜蜂图片预测,始终返回classes列表第0位的“ants”,而非正确的“bees”。

训练代码

from __future__ import print_function, division

import torch
import torch.nn as nn
import torch.optim as optim
from torch.optim import lr_scheduler
import torch.backends.cudnn as cudnn
import numpy as np
import torchvision
from torchvision import datasets, models, transforms
import matplotlib.pyplot as plt
import time
import os
import copy
import pickle


def main():
    data_transforms = {
        'train': transforms.Compose([
            transforms.RandomResizedCrop(224),
            transforms.RandomHorizontalFlip(),
            transforms.ToTensor(),
            transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
        ]),
        'val': transforms.Compose([
            transforms.Resize(256),
            transforms.CenterCrop(224),
            transforms.ToTensor(),
            transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
        ]),
    }

    data_dir = 'hymenoptera_data'
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir, x),
                                              data_transforms[x])
                      for x in ['train', 'val']}
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=4,
                                                  shuffle=True, num_workers=4)
                   for x in ['train', 'val']}
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'val']}
    class_names = image_datasets['train'].classes

    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")

    def imshow(inp, title=None):
        """Imshow for Tensor."""
        inp = inp.numpy().transpose((1, 2, 0))
        mean = np.array([0.485, 0.456, 0.406])
        std = np.array([0.229, 0.224, 0.225])
        inp = std * inp + mean
        inp = np.clip(inp, 0, 1)
        plt.imshow(inp)
        if title is not None:
            plt.title(title)
        plt.pause(0.001)  # pause a bit so that plots are updated

    # Get a batch of training data
    inputs, classes = next(iter(dataloaders['train']))

    # Make a grid from batch
    out = torchvision.utils.make_grid(inputs)

    imshow(out, title=[class_names[x] for x in classes])

    def train_model(model, criterion, optimizer, scheduler, num_epochs=25):
        since = time.time()

        best_model_wts = copy.deepcopy(model.state_dict())
        best_acc = 0.0

        for epoch in range(num_epochs):
            print(f'Epoch {epoch}/{num_epochs - 1}')
            print('-' * 10)

            # Each epoch has a training and validation phase
            for phase in ['train', 'val']:
                if phase == 'train':
                    model.train()  # Set model to training mode
                else:
                    model.eval()  # Set model to evaluate mode

                running_loss = 0.0
                running_corrects = 0

                # Iterate over data.
                for inputs, labels in dataloaders[phase]:
                    inputs = inputs.to(device)
                    labels = labels.to(device)

                    # zero the parameter gradients
                    optimizer.zero_grad()

                    # forward
                    # track history if only in train
                    with torch.set_grad_enabled(phase == 'train'):
                        outputs = model(inputs)
                        _, preds = torch.max(outputs, 1)
                        loss = criterion(outputs, labels)

                        # backward + optimize only if in training phase
                        if phase == 'train':
                            loss.backward()
                            optimizer.step()

                    # statistics
                    running_loss += loss.item() * inputs.size(0)
                    running_corrects += torch.sum(preds == labels.data)
                if phase == 'train':
                    scheduler.step()

                epoch_loss = running_loss / dataset_sizes[phase]
                epoch_acc = running_corrects.double() / dataset_sizes[phase]

                print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')

                # deep copy the model
                if phase == 'val' and epoch_acc > best_acc:
                    best_acc = epoch_acc
                    best_model_wts = copy.deepcopy(model.state_dict())

            print()

        time_elapsed = time.time() - since
        print(f'Training complete in {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s')
        print(f'Best val Acc: {best_acc:4f}')

        # load best model weights
        model.load_state_dict(best_model_wts)
        return model

    def visualize_model(model, num_images=6):
        was_training = model.training
        model.eval()
        images_so_far = 0
        fig = plt.figure()

        with torch.no_grad():
            for i, (inputs, labels) in enumerate(dataloaders['val']):
                inputs = inputs.to(device)
                labels = labels.to(device)

                outputs = model(inputs)
                _, preds = torch.max(outputs, 1)

                for j in range(inputs.size()[0]):
                    images_so_far += 1
                    ax = plt.subplot(num_images // 2, 2, images_so_far)
                    ax.axis('off')
                    ax.set_title(f'predicted: {class_names[preds[j]]}')
                    imshow(inputs.cpu().data[j])

                    if images_so_far == num_images:
                        model.train(mode=was_training)
                        return
            model.train(mode=was_training)

    model_ft = models.resnet18(pretrained=True)
    num_ftrs = model_ft.fc.in_features
    # Here the size of each output sample is set to 2.
    # Alternatively, it can be generalized to nn.Linear(num_ftrs, len(class_names)).
    model_ft.fc = nn.Linear(num_ftrs, 2)

    model_ft = model_ft.to(device)

    criterion = nn.CrossEntropyLoss()

    # Observe that all parameters are being optimized
    optimizer_ft = optim.SGD(model_ft.parameters(), lr=0.001, momentum=0.9)

    # Decay LR by a factor of 0.1 every 7 epochs
    exp_lr_scheduler = lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)

    ###
    # save using pickle

    # pickle.dump(model_ft, open('model.pkl', 'wb'))

    ###
    # save using torch
    # def save_model(model, best_acc):
    #     state = {
    #         'model': model_ft,
    #         'acc': best_acc,
    #     }

    torch.save(model_ft, './best_model.pth')



if __name__ == '__main__':
    main()

预测代码

# to be worked on

from __future__ import print_function, division

import torch

import numpy as np

from torchvision import transforms

import PIL.Image as Image

classes = [

    "ants",
    "bees",
]

# loading model
model = torch.load('best_model.pth')

# transform the image
mean = np.array([0.485, 0.456, 0.406])
std = np.array([0.229, 0.224, 0.225])

image_transforms = transforms.Compose([
    transforms.Resize((224, 224,)),
    transforms.ToTensor(),
    transforms.Normalize(torch.Tensor(mean), torch.Tensor(std))

])


def classify(model, image_transforms, image_path, classes):
    model = model.eval()
    image = Image.open(image_path)
    image = image_transforms(image).float()
    image = image.unsqueeze(0)

    output = model(image)
    _, predicted = torch.max(output.data, 1)

    print(classes[predicted.item()])

classify(model,image_transforms,"beeimage.jpg",classes)

预期与实际输出

预期输出:bees
实际输出:ants

实际运行日志:

C:\Users\prasa\Desktop\DL\venv\Scripts\python.exe C:\Users\prasa\Desktop\DL\callmod1.py 
ants

Process finished with exit code 0

问题根源及修复方案

1. 模型未训练直接保存

训练代码中仅定义了train_model训练函数,但从未调用该函数,直接保存了初始化状态的模型。此时模型的全连接层参数是随机值,预测结果完全随机,大概率偏向列表首个类别。

修复:在保存模型前调用训练函数,替换原有保存代码:

# 调用训练函数得到训练好的模型
model_ft = train_model(model_ft, criterion, optimizer_ft, exp_lr_scheduler, num_epochs=25)
# 保存训练后的最佳模型
torch.save(model_ft, './best_model.pth')

2. 图像预处理不一致

训练时验证集的预处理流程是Resize(256)→CenterCrop(224),但预测时直接Resize((224,224)),会导致图像比例失真,干扰模型识别。

修复:统一预测与训练的预处理流程:

image_transforms = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(torch.Tensor(mean), torch.Tensor(std))
])

3. 设备不匹配

训练时模型可能运行在GPU(cuda)上,预测时如果在CPU环境,模型与输入张量的设备不一致会导致隐性错误。

修复:在预测函数中添加设备匹配逻辑,并关闭梯度计算节省资源:

def classify(model, image_transforms, image_path, classes):
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    model = model.to(device).eval()
    image = Image.open(image_path)
    image = image_transforms(image).float()
    image = image.unsqueeze(0).to(device)  # 将图像张量转移到对应设备

    with torch.no_grad():  # 推理阶段关闭梯度计算
        output = model(image)
        _, predicted = torch.max(output.data, 1)

    print(classes[predicted.item()])

4. 类别顺序可能不一致

训练时class_names由数据集文件夹自动生成,顺序依赖文件夹名称排序;手动定义的classes = ["ants", "bees"]可能与训练时的顺序不符。

修复:训练时保存类别顺序,预测时加载:
训练代码中添加保存逻辑:

# 保存训练时的类别顺序
with open('class_names.pkl', 'wb') as f:
    pickle.dump(class_names, f)

预测代码中加载类别顺序:

import pickle
# 加载训练时的类别顺序
with open('class_names.pkl', 'rb') as f:
    classes = pickle.load(f)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 09:05:14