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

PyTorch模型与权重保存的区别及推理预测异常问题求助

PyTorch中保存模型与保存权重的区别
  • 保存完整模型:使用torch.save(model, 'model.pt'),保存内容包含模型结构、所有可学习权重参数,甚至可附带优化器状态。加载时直接通过model = torch.load('model.pt')即可使用,无需重新定义模型结构。缺点是文件体积较大,且依赖特定PyTorch版本与模型定义环境,跨环境加载可能失败。
  • 保存权重(state_dict):使用torch.save(model.state_dict(), 'weights.pt'),仅保存模型的可学习参数(权重、偏置等),文件体积小,兼容性强。加载时需先手动定义好与训练时完全一致的模型结构,再通过model.load_state_dict(torch.load('weights.pt'))加载参数。这是PyTorch官方推荐的方式,灵活性更高。
推理代码问题排查与修正

你的训练测试准确率达95%但推理全错,核心原因是输入数据的格式/预处理逻辑与训练阶段不一致,结合代码细节,修正点如下:

1. 图像通道格式不匹配

cv2.imread读取的图像为BGR格式,而你使用的WideResNet预训练模型的训练数据基于RGB格式。若训练时用PIL库(Image.open())读取图像(默认RGB),推理时的BGR输入会导致模型输入分布完全偏离训练数据,直接引发预测错误。
修正:添加BGR转RGB的步骤:

x = cv2.imread('anormal_9979_AVsp2000_ciclo6.png')
x = cv2.cvtColor(x, cv2.COLOR_BGR2RGB)  # 新增该行

2. 预处理顺序优化(更符合常规流程)

T.Resize放在T.ToTensor()之前更高效,因为对numpy数组(cv2读取的图像格式)或PIL图像做Resize,比处理Tensor更节省资源:

transform = T.Compose([
    T.ToPILImage(),  # 将cv2读取的numpy数组转为PIL图像
    T.Resize((64, 64)),
    T.ToTensor(),
    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

3. 移除冗余数据类型转换

代码中先将Tensor转为double()再转回float()完全多余,直接保留ToTensor()默认的float32格式即可:

batch_t = torch.unsqueeze(x_t, 0)  # 去掉.double()

4. 删除无用代码

torch.nn.CrossEntropyLoss()未赋值给变量也未使用,属于冗余代码,直接删除。

5. 确认模型加载逻辑正确性

  • 若训练时用torch.save(resnet.state_dict(), 'wide_resnet101_2.pt')保存权重,当前加载方式正确;若训练时保存的是完整模型(torch.save(resnet, 'wide_resnet101_2.pt')),则需改为resnet = torch.load('wide_resnet101_2.pt', map_location=torch.device('cpu')),无需重新定义模型结构与修改fc层。
  • 确保训练时修改的fc层与推理时完全一致:训练时是否同样将fc改为torch.nn.Linear(resnet.fc.in_features, 2)?若fc输出维度不一致,也会导致预测错误。

修正后的完整推理代码

import cv2
import torch
import torchvision.transforms as T
import torchvision.models as models

x = cv2.imread('anormal_9979_AVsp2000_ciclo6.png')
x = cv2.cvtColor(x, cv2.COLOR_BGR2RGB)

transform = T.Compose([
    T.ToPILImage(),
    T.Resize((64, 64)),
    T.ToTensor(),
    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

x_t = transform(x)
batch_t = torch.unsqueeze(x_t, 0)

# 初始化模型并加载权重
resnet = models.wide_resnet101_2(pretrained=False)
resnet.fc = torch.nn.Linear(resnet.fc.in_features, 2)
resnet.load_state_dict(torch.load('wide_resnet101_2.pt', map_location=torch.device('cpu')))
resnet.eval()

with torch.no_grad():  # 推理阶段关闭梯度计算,节省内存
    net = resnet(batch_t)
    predictions = torch.argmax(net, axis=1)
    print(predictions)

关于model.pt是否包含权重的说明

  • 若用torch.save(model, 'model.pt')保存完整模型,文件包含模型结构+所有权重参数;
  • 若用torch.save(model.state_dict(), 'model.pt')保存,文件仅包含权重参数,无模型结构。
    你当前加载的是state_dict,因此需要先定义模型结构再加载权重。

内容的提问来源于stack exchange,提问作者Raúl Nieto

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 05:47:40