如何将nn.DataParallel训练的ResNet-50 checkpoint加载到非并行模型?
解决方案
核心问题分析
你遇到的问题根源有两个:
- 用
nn.DataParallel训练的模型,state_dict里所有参数键都带module.前缀,而单GPU模型的参数键没有这个前缀,直接加载会出现键不匹配。 - 训练时的模型可能包含额外的
aux_bn层参数,但你当前加载的标准ResNet-50没有这部分结构,导致出现"意外键"错误。
具体解决代码
直接用下面的代码处理checkpoint,同时解决前缀和多余键的问题:
import torch from torchvision.models import resnet50 # 初始化单GPU版本的ResNet-50(无需DataParallel包装) model = resnet50() # 加载保存的checkpoint文件 checkpoint = torch.load("你的checkpoint路径.pth") # 注意:有些checkpoint会把state_dict存在特定键下(比如'net'/'model'),需要根据实际调整 state_dict = checkpoint.get("state_dict", checkpoint) # 处理state_dict:移除module.前缀 + 过滤模型不存在的键(比如aux_bn相关) processed_state_dict = {} for key, value in state_dict.items(): # 去掉module.前缀 if key.startswith("module."): key = key[7:] # 只保留当前模型存在的参数键 if key in model.state_dict(): processed_state_dict[key] = value # 加载处理后的参数,strict=False可以兼容少量未覆盖的参数(可选,这里已经过滤的话也可以用strict=True) model.load_state_dict(processed_state_dict, strict=False)
额外排查技巧
如果还是有问题,可以先打印对比两边的参数键,明确不匹配的部分:
# 打印原checkpoint的前10个键 print("Checkpoint中的参数键:", list(state_dict.keys())[:10]) # 打印当前模型的前10个键 print("当前模型的参数键:", list(model.state_dict().keys())[:10])
这样就能精准看到哪些键需要调整或过滤。
注意事项
- 确保你加载的模型结构和训练时的模型主体一致(除了DataParallel和aux_bn部分),如果训练时是自定义的ResNet变种,可能需要同步模型结构,或者进一步过滤不匹配的键。
- 加载checkpoint时,如果是在CPU环境加载GPU保存的模型,需要加上
map_location=torch.device('cpu')参数,避免设备不匹配错误。
内容的提问来源于stack exchange,提问作者Fatemeh
相关产品推荐
相关产品推荐

