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

PyTorch加载已保存模型时重复训练问题及解决咨询

正确加载PyTorch已训练模型的方法

问题根源

你遇到的不是模型加载失败,而是加载模型后误执行了训练代码,导致参数被重新训练覆盖。另外两种加载方式本身存在细节问题:

  • 方式一:torch.load(PATH)加载的是你保存的模型state_dict(字典格式的参数集合),不是完整的模型实例,无法直接用于推理。
  • 方式二:是正确的state_dict加载流程,但如果加载后又运行了train()循环,自然会重新训练。

正确加载流程

1. 确保加载环境有一致的模型定义

在加载模型的文件中,必须先定义和训练时完全相同的Net类:

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

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(1, 10, kernel_size=5)
        self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
        self.conv2_drop = nn.Dropout2d()
        self.fc1 = nn.Linear(320, 50)
        self.fc2 = nn.Linear(50, 10)

    def forward(self, x):
        x = F.relu(F.max_pool2d(self.conv1(x), 2))
        x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 2))
        x = x.view(-1, 320)
        x = F.relu(self.fc1(x))
        x = F.dropout(x, training=self.training)
        x = self.fc2(x)
        return F.log_softmax(x, dim=1)  # 补充dim参数避免警告

2. 加载参数并切换到评估模式

使用修正后的方式二加载,加载后不要调用训练函数,直接切换到评估模式:

PATH = "results/model.pth"
# 初始化模型实例
model = Net()
# 加载训练好的参数
model.load_state_dict(torch.load(PATH))
# 切换到评估模式(关闭dropout、batchnorm等训练专属行为)
model.eval()

3. 验证加载效果

直接运行测试代码验证模型性能,确认参数已正确加载:

# 定义测试函数(与训练时一致)
def test(model, test_loader):
    test_loss = 0
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            output = model(data)
            test_loss += F.nll_loss(output, target, size_average=False).item()
            pred = output.data.max(1, keepdim=True)[1]
            correct += pred.eq(target.data.view_as(pred)).sum()
    test_loss /= len(test_loader.dataset)
    print('\nTest set: Avg. loss: {:.4f}, Accuracy: {}/{} ({:.0f}%)\n'.format(
        test_loss, correct, len(test_loader.dataset),
        100. * correct / len(test_loader.dataset)))

# 准备测试数据集(与训练时一致)
import torchvision
test_loader = torch.utils.data.DataLoader(
  torchvision.datasets.MNIST('./files', train=False, download=True,
                             transform=torchvision.transforms.Compose([
                               torchvision.transforms.ToTensor(),
                               torchvision.transforms.Normalize(
                                 (0.1307,), (0.3081,))
                             ])),
  batch_size=1000, shuffle=True)

# 运行测试
test(model, test_loader)

额外注意事项

  • 推荐始终保存state_dict而非完整模型:torch.save(network.state_dict(), path),这种方式兼容性更强,不受环境、版本影响。如果要直接加载完整模型,需用torch.save(network, path),但不推荐。
  • 加载后必须调用model.eval(),否则dropout等层会继续随机丢弃神经元,导致推理结果不稳定。
  • 模型类定义必须和训练时完全一致,哪怕是forward函数的微小修改(比如缺少dim参数)都可能导致加载失败或性能异常。

内容的提问来源于stack exchange,提问作者Timothy Oh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 08:40:42