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
相关产品推荐
相关产品推荐

