验证预训练模型访问
PyTorch - 访问预训练模型
Section titled “PyTorch - 访问预训练模型”在使用预训练模型之前,我们先验证一下能否访问和加载 PyTorch 生态系统提供的模型,特别是包含流行计算机视觉模型的 torchvision 库。
与旧的 Caffe2 方法不同(模型通常是安装目录中的特定文件),PyTorch 通常在您首次请求预训练模型时按需下载预训练权重。这些权重通常缓存在一个中心目录中(例如,在 Linux/macOS 上是 ~/.cache/torch/hub/checkpoints/)。
您可以使用 torchvision.models 轻松检查可用模型并加载一个。我们来尝试加载 SqueezeNet 1.1 模型,它在 Caffe2 中也是可用的。
首先,确保您已经安装了 torch 和 torchvision(参见安装章节)。然后,运行以下 Python 脚本:
import torchimport 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 来类似地加载。
现在您已经准备好使用预训练模型执行实际任务,例如图像分类。