使用预训练模型进行图像分类
PyTorch - 使用预训练模型进行图像分类
Section titled “PyTorch - 使用预训练模型进行图像分类”在本课中,您将学习如何使用 torchvision 中的预训练深度学习模型对图像中的对象进行分类。我们将使用在大型 ImageNet 数据集上训练的 SqueezeNet 1.1(类似于原始教程)或 ResNet 模型。
我们将涵盖以下步骤:加载图像、对其进行预处理 (preprocessing) 以匹配模型的输入要求、加载预训练模型、执行推理 (inference) 并解释结果。
首先,导入必要的 Python 包:
import torchimport torchvision.models as modelsimport torchvision.transforms as transformsfrom PIL import Image # Pillow 库用于加载图像import requestsimport numpy as npimport jsonimport os确保您已安装 Pillow (pip install Pillow) 和 requests (pip install requests)。
设置模型和预处理
Section titled “设置模型和预处理”预训练模型期望输入图像以特定的方式进行处理(例如,调整大小、裁剪、归一化)。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.")加载和处理图像
Section titled “加载和处理图像”现在,加载一个示例图像。您可以下载一个,或者使用本地文件。我们将使用一个名为 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)信息量不大。我们需要将这个索引映射到人类可读的类别名称。
将索引映射到类别名称
Section titled “将索引映射到类别名称”我们可以下载标准的 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})")