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

如何使用.pth格式的预训练PyTorch模型在独立Python文件中进行预测

PyTorch .pth预训练模型预测指南

以下是使用.pth格式预训练模型对新数据进行预测的完整实操步骤:

1. 准备环境与依赖

  • 确保已安装PyTorch及对应业务依赖(图像任务需torchvision,文本任务需transformers等),基础安装命令:
    pip install torch torchvision
    
  • 提前确认模型的网络结构定义:.pth文件仅存储模型权重,必须使用与训练时完全一致的网络结构才能成功加载。

2. 定义模型结构

情况1:自定义模型

直接复制训练阶段的模型类代码到当前脚本,例如一个简单的图像分类CNN:

import torch.nn as nn
import torch.nn.functional as F

class CustomCNN(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 16, 3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(16, 32, 3, padding=1)
        self.fc1 = nn.Linear(32 * 8 * 8, 128)
        self.fc2 = nn.Linear(128, num_classes)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 32 * 8 * 8)
        x = F.relu(self.fc1(x))
        x = self.fc2(x)
        return x

情况2:TorchVision官方预训练模型(如ResNet、VGG)

若训练时未修改模型结构,直接调用即可;若调整过最后一层(如自定义类别数),需手动修改:

from torchvision import models

# 未修改结构的情况
model = models.resnet50(pretrained=False)  # 先不加载官方预训练权重

# 修改最后一层的情况(假设训练时将类别数改为20)
num_classes = 20
model.fc = nn.Linear(model.fc.in_features, num_classes)

3. 加载预训练权重

初始化模型后,加载.pth权重并设置运行设备(CPU/GPU):

import torch

# 自动选择设备:优先GPU,否则CPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 初始化模型并移至目标设备
model = CustomCNN(num_classes=10).to(device)  # 或上述TorchVision模型

# 加载权重:若训练用GPU、当前用CPU,需加map_location='cpu'
model.load_state_dict(torch.load("your_model.pth", map_location=device))

# 切换至预测模式:关闭Dropout、BatchNorm的训练行为
model.eval()

4. 预处理新数据

预处理流程必须与训练时完全一致,否则会导致预测结果失真。以图像数据为例:

from PIL import Image
from torchvision import transforms

# 复刻训练时的预处理流程
transform = transforms.Compose([
    transforms.Resize((32, 32)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 加载单张测试图并预处理
img = Image.open("test_image.jpg").convert("RGB")
input_tensor = transform(img)
# 添加batch维度:模型默认接受批量输入,需从[C, H, W]转为[1, C, H, W]
input_batch = input_tensor.unsqueeze(0).to(device)

5. 执行预测并解析结果

关闭梯度计算以节省资源,然后完成预测并解析输出:

# 关闭梯度计算,避免不必要的内存占用
with torch.no_grad():
    outputs = model(input_batch)

# 分类任务结果解析示例
# 方法1:获取预测类别索引
predicted_idx = torch.argmax(outputs, dim=1).item()

# 方法2:获取每个类别的概率
probabilities = F.softmax(outputs, dim=1).squeeze().tolist()

# 若有类别标签映射,可转换为类别名称
class_names = ["cat", "dog", "bird", ...]  # 需与训练时的标签顺序一致
predicted_class = class_names[predicted_idx]

print(f"预测类别:{predicted_class}")
print(f"类别概率:{probabilities}")

常见问题排查

  • 模型结构不匹配报错:检查当前模型的层数、输入输出维度是否与训练时完全一致,重点核对自定义的全连接层、卷积层参数。
  • 设备不兼容:训练用GPU、当前用CPU时,需在torch.load中添加map_location='cpu'。
  • 预测结果不稳定:确保已调用model.eval(),否则Dropout、BatchNorm会在预测时仍使用训练模式,导致结果波动。
  • 预处理不一致:核对图像resize尺寸、归一化均值方差、是否转RGB等细节,必须与训练代码完全相同。

内容的提问来源于stack exchange,提问作者Ammar Ahmed Siddiqui

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 01:09:25