Skip to content

验证预训练模型访问

在使用预训练模型之前,我们先验证一下能否访问和加载 PyTorch 生态系统提供的模型,特别是包含流行计算机视觉模型的 torchvision 库。

与旧的 Caffe2 方法不同(模型通常是安装目录中的特定文件),PyTorch 通常在您首次请求预训练模型时按需下载预训练权重。这些权重通常缓存在一个中心目录中(例如,在 Linux/macOS 上是 ~/.cache/torch/hub/checkpoints/)。

您可以使用 torchvision.models 轻松检查可用模型并加载一个。我们来尝试加载 SqueezeNet 1.1 模型,它在 Caffe2 中也是可用的。

首先,确保您已经安装了 torch 和 torchvision(参见安装章节)。然后,运行以下 Python 脚本:

import torch
import torchvision.models as models
# List some available models (optional)
# print(dir(models))
print("Attempting to load pre-trained SqueezeNet 1.1 model...")
try:
# Load the SqueezeNet 1.1 architecture with pre-trained weights
# PyTorch will download weights if not already cached
model = models.squeezenet1_1(pretrained=True)
# Set the model to evaluation mode (important for inference)
model.eval()
print("Successfully loaded pre-trained SqueezeNet 1.1 model.")
print("Model structure:")
# Print the model structure (optional, can be very long)
# print(model)
print("Model ready for inference.")
except ImportError:
print("Error: torchvision is not installed. Please install it using: pip install torchvision")
except Exception as e:
print(f"An error occurred: {e}")
print("Please ensure you have an internet connection for the first download.")

如果脚本成功运行,您将看到确认模型已加载的输出。首次运行时,您可能会看到下载进度。

Attempting to load pre-trained SqueezeNet 1.1 model...
Successfully loaded pre-trained SqueezeNet 1.1 model.
Model structure:
Model ready for inference.

这确认您可以从 torchvision 访问和加载预训练模型。许多其他模型,如 ResNet、VGG、MobileNet 等,也可以通过在 torchvision.models 中调用各自的函数并设置 pretrained=True 来类似地加载。

现在您已经准备好使用预训练模型执行实际任务,例如图像分类。