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

如何将nn.DataParallel训练的ResNet-50 checkpoint加载到非并行模型?

解决方案

核心问题分析

你遇到的问题根源有两个:

  1. 用nn.DataParallel训练的模型,state_dict里所有参数键都带module.前缀,而单GPU模型的参数键没有这个前缀,直接加载会出现键不匹配。
  2. 训练时的模型可能包含额外的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 02:06:09