torch.load后模型性能骤降,评估后回升的问题求助
问题描述
训练模型后保存权重,即时测试准确率为0.95;加载权重后首次测试准确率接近0(随机猜测水平),用全测试集评估后准确率回升至0.8,但仍有性能损失。已确认model.state_dict()在评估前后完全一致。
模型代码:
class MyModel(nn.Module): def __init__(self, feat_dim, num_classes): super(MyModel, self).__init__() self.model_resnet = models.resnet50(pretrained=False) num_ftrs = self.model_resnet.fc.in_features self.model_resnet.fc = nn.Identity() self.head1 = nn.Sequential( nn.Linear(num_ftrs, num_ftrs), nn.ReLU(inplace=True), nn.Linear(num_ftrs, feat_dim) ) self.head2 = nn.Linear(num_ftrs, num_classes) def forward(self, x): self.eps=self.eps+1 x = self.model_resnet(x) feat = F.normalize(self.head1(x), dim=1) classes = self.head2(x) return feat,classes
保存与加载代码:
torch.save(model.state_dict(),"./test.pth") model.load_state_dict(torch.load("test.pth"))
问题原因与解决方案
核心原因1:模型训练/评估模式未切换
训练后即时测试时,你大概率调用了model.eval()切换到评估模式,此时ResNet中的BatchNorm层会使用训练阶段保存的running_mean和running_var,保证输出稳定,因此准确率达到0.95。但加载模型后测试时,你可能未切换到评估模式,模型仍处于train()状态:
- 训练模式下,BatchNorm会用当前输入batch的均值和方差做归一化,小batch的统计量偏差极大,导致输出完全随机,准确率接近0;
- 跑完整测试集后,训练模式下的BatchNorm会逐步更新
running_mean和running_var,使其接近真实分布,因此准确率回升,但统计量无法完全恢复到训练结束时的状态,所以存在性能损失。
核心原因2:非法实例变量self.eps
forward方法中直接修改self.eps,但__init__里未初始化该变量,首次调用forward会触发AttributeError。即使你在其他地方初始化了它,这个变量不属于模型的参数或缓冲区,不会被保存到state_dict中:加载模型后self.eps会被重置为初始值,若代码中实际用到该变量(当前代码未体现),会导致逻辑不一致。
解决方案
- 强制切换模型模式:加载模型后必须调用
model.eval(),并在测试时禁用梯度计算:
model.load_state_dict(torch.load("test.pth")) model.eval() # 切换到评估模式 with torch.no_grad(): # 执行测试逻辑 pass
处理
self.eps变量:- 如果是无用代码,直接删除
self.eps=self.eps+1这一行; - 如果确实需要该变量,在
__init__中初始化,若需随模型保存则注册为缓冲区:def __init__(self, feat_dim, num_classes): super(MyModel, self).__init__() # 若无需保存状态 self.eps = 0 # 若需随模型保存状态 self.register_buffer('eps', torch.tensor(0)) # 其他初始化代码...
- 如果是无用代码,直接删除
验证BatchNorm状态:检查加载后的模型中,ResNet各BatchNorm层的
running_mean和running_var是否与保存前一致,确保缓冲区被正确加载。
内容的提问来源于stack exchange,提问作者MarsEclipse
相关产品推荐
相关产品推荐

