Skip to content

使用预训练模型进行图像分类

PyTorch - 使用预训练模型进行图像分类

Section titled “PyTorch - 使用预训练模型进行图像分类”

在本课中,您将学习如何使用 torchvision 中的预训练深度学习模型对图像中的对象进行分类。我们将使用在大型 ImageNet 数据集上训练的 SqueezeNet 1.1(类似于原始教程)或 ResNet 模型。

我们将涵盖以下步骤:加载图像、对其进行预处理 (preprocessing) 以匹配模型的输入要求、加载预训练模型、执行推理 (inference) 并解释结果。

首先,导入必要的 Python 包:

import torch
import torchvision.models as models
import torchvision.transforms as transforms
from PIL import Image # Pillow 库用于加载图像
import requests
import numpy as np
import json
import os

确保您已安装 Pillow (pip install Pillow) 和 requests (pip install requests)。

预训练模型期望输入图像以特定的方式进行处理(例如,调整大小、裁剪、归一化)。torchvision.transforms 提供了执行此操作的工具。

# Load a pre-trained model (e.g., SqueezeNet 1.1 or ResNet18)
# 加载预训练模型 (例如 SqueezeNet 1.1 或 ResNet18)
# model = models.squeezenet1_1(pretrained=True)
model = models.resnet18(pretrained=True)
model.eval() # 将模型设置为评估模式(禁用 dropout 等)
# Define the image transformations
# 定义图像转换
# Models trained on ImageNet usually expect 224x224 images
# 在 ImageNet 上训练的模型通常期望输入 224x224 的图像
# Normalization values are standard for ImageNet
# 归一化值是 ImageNet 的标准值
preprocess = transforms.Compose([
transforms.Resize(256), # 将较短边调整为 256
transforms.CenterCrop(224), # 从中心裁剪 224x224
transforms.ToTensor(), # 将图像转换为 PyTorch Tensor(缩放到 [0, 1])
transforms.Normalize( # 使用 ImageNet 的均值和标准差进行归一化
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
print("Model and preprocessing pipeline ready.")

现在,加载一个示例图像。您可以下载一个,或者使用本地文件。我们将使用一个名为 image.jpg 的本地文件,将其放在与脚本相同的目录中(或者提供完整路径)。

image_path = 'image.jpg' # 确保此图像文件存在
if not os.path.exists(image_path):
print(f"Error: Image file not found at {image_path}")
print("Please download an image and save it as image.jpg or update the path.")
# Example download (requires internet)
# 示例下载(需要网络)
# try:
# url = "https://raw.githubusercontent.com/pytorch/hub/master/images/dog.jpg"
# img_data = requests.get(url).content
# with open(image_path, 'wb') as handler:
# handler.write(img_data)
# print(f"Downloaded example image to {image_path}")
# except Exception as e:
# print(f"Could not download example image: {e}")
# exit()
else:
print(f"Loading image from {image_path}")
try:
input_image = Image.open(image_path).convert('RGB') # 确保图像为 RGB 格式
except FileNotFoundError:
print(f"Error: Cannot open image file at {image_path}")
exit()
# Apply the preprocessing transformations
# 应用预处理转换
input_tensor = preprocess(input_image)
# Create a mini-batch (add batch dimension)
# 创建一个 mini-batch (添加批次维度)
# Models expect input shape [batch_size, channels, height, width]
# 模型期望的输入形状为 [批次大小, 通道数, 高度, 宽度]
input_batch = input_tensor.unsqueeze(0)
print("Image loaded and preprocessed.")
print("Input tensor shape:", input_batch.shape) # 应为 [1, 3, 224, 224]

现在,将处理后的图像张量 (tensor) 输入到模型中。我们使用 torch.no_grad(),因为我们不进行训练,这可以节省内存和计算。

with torch.no_grad(): # 禁用梯度计算
output = model(input_batch)
# The output contains raw scores (logits) for each class
# 输出包含每个类别的原始得分(logits)
print("Inference complete.")
print("Output tensor shape:", output.shape) # 对于 ImageNet 模型,形状为 [1, 1000]

模型的输出是 1000 个 ImageNet 类别的原始得分(logits)张量。为了获得概率 (probabilities),我们应用 Softmax 函数。然后,我们找到概率最高的类别。

# Apply Softmax to get probabilities
# 应用 Softmax 获取概率
probabilities = torch.nn.functional.softmax(output[0], dim=0)
# Get the top prediction
# 获取最高预测
top1_prob, top1_idx = torch.topk(probabilities, 1)
predicted_idx = top1_idx[0].item()
confidence = top1_prob[0].item()
print(f"\nPredicted Index: {predicted_idx}")
print(f"Confidence: {confidence:.4f}")

索引(例如,ImageNet 中表示“golden retriever”的 207)信息量不大。我们需要将这个索引映射到人类可读的类别名称。

我们可以下载标准的 ImageNet 类别索引映射。

# Download the ImageNet class index mapping
# 下载 ImageNet 类别索引映射
imagenet_classes_url = "https://raw.githubusercontent.com/pytorch/examples/main/imagenet/imagenet_classes.txt"
# Alternative: Use a cached JSON mapping if available
# 备选项:如果可用,使用缓存的 JSON 映射
# imagenet_classes_url = "https://s3.amazonaws.com/deep-learning-models/image-models/imagenet_class_index.json"
try:
response = requests.get(imagenet_classes_url)
response.raise_for_status() # 如果状态码异常则抛出异常
# Simple text format: each line is a class name
# 简单文本格式:每行是一个类别名称
labels = response.text.split('\n')
# Remove empty lines if any
# 移除空行(如果存在)
labels = [label for label in labels if label]
if predicted_idx < len(labels):
predicted_class_name = labels[predicted_idx]
print(f"Predicted Class Name: {predicted_class_name}")
else:
print(f"Error: Predicted index {predicted_idx} is out of bounds for the loaded labels (length {len(labels)}). Using index directly.")
predicted_class_name = f"Class Index {predicted_idx}"
except requests.exceptions.RequestException as e:
print(f"Could not download or process class labels: {e}")
predicted_class_name = f"Class Index {predicted_idx}"
print(f"\nFinal Prediction: {predicted_class_name} (Confidence: {confidence:.4f})")