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

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会被重置为初始值,若代码中实际用到该变量(当前代码未体现),会导致逻辑不一致。

解决方案

  1. 强制切换模型模式:加载模型后必须调用model.eval(),并在测试时禁用梯度计算:
model.load_state_dict(torch.load("test.pth"))
model.eval()  # 切换到评估模式
with torch.no_grad():
    # 执行测试逻辑
    pass
  1. 处理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))
          # 其他初始化代码...
      
  2. 验证BatchNorm状态:检查加载后的模型中,ResNet各BatchNorm层的running_mean和running_var是否与保存前一致,确保缓冲区被正确加载。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 12:31:26