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

PyTorch加载RSP-ResNet50预训练权重时输出不一致问题求助

解决方案

1. 正确提取模型权重

你加载的checkpoint文件里,模型权重存在model键下,不是常规的state_dict,先提取这部分核心权重:

res50_state = torch.load("rsp-resnet-50-ckpt.pth")["model"]

2. 处理权重键名前缀(关键步骤)

这个仓库保存的权重键大概率带model.前缀(比如model.conv1.weight),但你初始化的ResNet模型参数键是不带该前缀的(比如conv1.weight),必须去掉前缀才能匹配:

# 遍历权重字典,移除键名开头的"model."
adjusted_state_dict = {k.replace("model.", ""): v for k, v in res50_state.items()}

3. 加载权重并验证

用处理后的权重字典加载,优先尝试strict=True(确保核心层权重完全匹配):

res50.load_state_dict(adjusted_state_dict, strict=True)

如果最后一层全连接层(fc)的类别数和你的任务不匹配,再改用strict=False跳过这部分不匹配的键即可。

4. 固定推理结果(解决输出不一致问题)

之前同一图像输出不同,核心原因是模型处于训练模式,Dropout、BatchNorm这类层会引入随机性。推理前必须切换到eval模式:

res50.eval()

同时关闭梯度计算,避免额外干扰:

with torch.no_grad():
    # 你的推理代码,比如 output = res50(input_img)

完整示例代码

import torch
from torchvision.models import resnet50

# 初始化ResNet50模型(注意不要用pretrained=True,我们要加载自定义权重)
res50 = resnet50(pretrained=False)
# 按需修改最后一层fc的类别数,比如你的场景分类有100类就改成100
# res50.fc = torch.nn.Linear(2048, 100)

# 加载并处理权重
ckpt = torch.load("rsp-resnet-50-ckpt.pth")
model_weights = ckpt["model"]
adjusted_weights = {k.replace("model.", ""): v for k, v in model_weights.items()}
res50.load_state_dict(adjusted_weights, strict=False)  # fc层不匹配时用False

# 切换到推理模式
res50.eval()

# 测试推理
input_img = torch.randn(1, 3, 224, 224)  # 模拟输入图像
with torch.no_grad():
    output = res50(input_img)
    pred_class = torch.argmax(output, dim=1)
print(pred_class)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 08:13:29