加载CycleGan训练得到的latest_net_G_A.pth模型报错如何解决
报错原因分析
- 第一个报错
AttributeError: ‘collections.OrderedDict’ has no attribute ‘eval’:你训练时只保存了模型的权重参数(state_dict,存储结构为OrderedDict),未保存完整模型对象,直接torch.load得到的是参数字典,不存在模型类才有的eval方法。 - 第二个报错
AttributeError : ‘torch.C_.Generator’ object has no attribute ‘load_state_dict’:你实例化的Generator不是继承自torch.nn.Module的自定义生成器类,大概率是导包错误引入了PyTorch内置的同名对象,这类内置对象没有load_state_dict方法。
解决步骤
- 先找到你训练CycleGAN时定义生成器的代码,确认生成器类的名称、初始化参数,必须保证加载时使用的生成器类和训练时完全一致。
- 按如下代码模板修改加载逻辑:
import torch # 替换为你自己的生成器类导入路径,不要导错其他同名类 from your_train_code_file import Generator # 1. 实例化生成器,初始化参数和训练时完全对齐,比如输入输出通道、残差块数量等 model = Generator(input_nc=3, output_nc=3, n_blocks=9) # 2. 加载权重字典,map_location参数适配你当前运行的设备 state_dict = torch.load("latest_net_G_A.pth", map_location=torch.device("cpu")) # 3. 权重载入模型 model.load_state_dict(state_dict) # 4. 切换为评估模式 model.eval()
常见兼容问题处理
如果出现权重key不匹配的报错,大概率是训练时使用了多卡分布式训练,权重key多了module.前缀,按如下方式处理后再加载:
from collections import OrderedDict new_state_dict = OrderedDict() for k, v in state_dict.items(): # 去掉key开头的'module.'前缀 if k.startswith("module."): name = k[7:] else: name = k new_state_dict[name] = v model.load_state_dict(new_state_dict)
内容的提问来源于stack exchange,提问作者Wilson Wii
相关产品推荐
相关产品推荐

