PyTorch加载检查点二次训练后测试出现State_dict错误求助
解决方案
1. 修正模型加载的DataParallel顺序与state_dict匹配问题
你的错误大概率是模型state_dict的键名不匹配导致的——第一次训练时模型是否用nn.DataParallel包装,直接决定了保存的state_dict键是否带有module.前缀,二次加载时的包装顺序如果和第一次不一致,就会出现键不匹配。
两种常见场景的修复方式:
场景A:第一次训练时,模型未用nn.DataParallel包装就保存了权重
此时保存的state_dict键没有module.前缀,二次加载时要先加载权重,再包装DataParallel:
# 1. 先创建原始模型实例 model = YourOriginalModelClass() # 2. 加载无module.前缀的模型权重 model.load_state_dict(torch.load(path_to_saved_model_weights, map_location=dev)) # 3. 再包装DataParallel(注意output_device是单个设备,不是列表) model = nn.DataParallel(model, device_ids=[0,1], output_device=0) model.to(dev)
场景B:第一次训练时,模型是用nn.DataParallel包装后保存的权重
此时保存的state_dict键带有module.前缀,两种修复方式:
- 方式1:先包装DataParallel再加载权重:
model = YourOriginalModelClass() model = nn.DataParallel(model, device_ids=[0,1], output_device=0) model.to(dev) # 直接加载带module.前缀的state_dict,键名匹配 model.load_state_dict(torch.load(path_to_saved_model_weights, map_location=dev))
- 方式2:手动移除state_dict中的
module.前缀,再加载到原始模型:
model = YourOriginalModelClass() state_dict = torch.load(path_to_saved_model_weights, map_location=dev) # 移除所有键名的module.前缀 state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()} model.load_state_dict(state_dict) model = nn.DataParallel(model, device_ids=[0,1], output_device=0) model.to(dev)
2. 修正优化器加载与学习率修改的顺序
你直接修改lr后加载优化器state_dict,会导致两个问题:一是加载的state_dict会覆盖你设置的新lr;二是如果模型参数结构不匹配(比如DataParallel导致的参数名变化),会触发state_dict错误。正确步骤是:
# 1. 先确保模型已经正确加载(完成上面的步骤) # 2. 创建优化器,先用原始训练时的lr(或任意值,后续会修改) optimizer = torch.optim.Adam(model.parameters(), lr=original_training_lr, weight_decay=2e-7) # 3. 加载优化器的state_dict optimizer.load_state_dict(torch.load(path_to_saved_optimizer_weights)) # 4. 最后修改学习率 for param_group in optimizer.param_groups: param_group['lr'] = your_new_lr
3. 额外注意事项
nn.DataParallel的output_device参数必须是单个设备(如0),不能传列表,这不符合PyTorch的API规范,可能引发潜在错误。- 加载权重前,确保两次训练的模型结构完全一致(包括层数、参数数量、命名),任何结构变更都会导致state_dict不匹配。
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

