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

