加载PyTorch Lightning训练的.pt模型遇ModuleNotFoundError,如何挽救?
问题描述
尝试加载一个用torch.save保存、经PyTorch Lightning训练的.pt模型,之前能正常加载评估,但现在加载时报错:
--------------------------------------------------------------------------- ModuleNotFoundError Traceback (most recent call last) ----> 6 torch.load(model.pt) File C:\ProgramData\Anaconda3\envs\Env1\lib\site-packages\torch\serialization.py:712, in load(f, map_location, pickle_module, **pickle_load_args) 710 opened_file.seek(orig_position) 711 return torch.jit.load(opened_file) --> 712 return _load(opened_zipfile, map_location, pickle_module, **pickle_load_args) 713 return _legacy_load(opened_file, map_location, pickle_module, **pickle_load_args) File C:\ProgramData\Anaconda3\envs\Env1\lib\site-packages\torch\serialization.py:1049, in _load(zip_file, map_location, pickle_module, pickle_file, **pickle_load_args) 1047 unpickler = UnpicklerWrapper(data_file, **pickle_load_args) 1048 unpickler.persistent_load = persistent_load --> 1049 result = unpickler.load() 1051 torch._utils._validate_loaded_sparse_tensors() 1053 return result File C:\ProgramData\Anaconda3\envs\Env1\lib\site-packages\torch\serialization.py:1042, in _load.<locals>.UnpicklerWrapper.find_class(self, mod_name, name) 1040 pass 1041 mod_name = load_module_mapping.get(mod_name, mod_name) --> 1042 return super().find_class(mod_name, name) ModuleNotFoundError: No module named 'pytorch_lightning.accelerators.cuda'
已了解到torch.save保存整个模块的方式存在缺陷,序列化数据和保存时的类结构、目录绑定,现在想挽救这个效果较好的模型,询问是否有方法在结构变更的情况下访问模型或其参数。
解决方法
以下几种方法可以尝试挽救模型:
方法1:临时映射缺失的模块路径
利用PyTorch加载时的load_module_mapping参数,把找不到的旧模块路径映射到当前环境存在的模块。比如PyTorch Lightning模块结构变更后,pytorch_lightning.accelerators.cuda可能被合并到其他模块,可手动配置映射:
import torch from pytorch_lightning.accelerators import Accelerator # 定义模块映射规则,根据当前PL版本调整 load_mapping = { 'pytorch_lightning.accelerators.cuda': 'pytorch_lightning.accelerators', # 若明确知道类的新路径,也可直接映射到类 # 'pytorch_lightning.accelerators.cuda.CUDAAccelerator': Accelerator } # 加载模型时传入映射 model = torch.load('model.pt', load_module_mapping=load_mapping)
若不确定新模块路径,可查看当前环境下pytorch_lightning.accelerators的结构,调整映射规则。
方法2:只加载模型参数,重新构建模型结构
如果能还原训练时的模型代码(或回忆起完整结构),可以避开pickle对类路径的依赖,只加载参数字典:
- 重新定义与训练时完全一致的PyTorch Lightning模型类(包括层结构、参数初始化等)
- 提取并加载模型参数:
import torch from your_module import YourLightningModel # 替换为你的模型类 # 初始化模型 model = YourLightningModel() # 尝试从保存文件中提取state_dict try: import pickle with open('model.pt', 'rb') as f: checkpoint = pickle.load(f, fix_imports=False, encoding='latin1') # Lightning模型参数通常存在'state_dict'键下,或直接是模型的state_dict if 'state_dict' in checkpoint: model.load_state_dict(checkpoint['state_dict']) else: model.load_state_dict(checkpoint.state_dict()) except: # 备选加载方式 model.load_state_dict(torch.load('model.pt', map_location='cpu')['state_dict'])
核心是不加载整个模型对象,只加载参数字典,只要模型结构匹配就能恢复模型功能。
方法3:降级PyTorch Lightning到训练时的版本
如果上述方法无效,可以安装训练模型时使用的PyTorch Lightning版本,让模块路径匹配:
# 替换为训练时的PL版本,例如1.7.0 pip install pytorch-lightning==1.7.0
加载成功后,立即将模型重新保存为state_dict格式,方便后续在新版本中使用:
model = torch.load('model.pt') torch.save(model.state_dict(), 'model_state_dict.pt')
方法4:手动修改pickle文件中的模块路径(进阶)
熟悉pickle结构的话,可以直接修改二进制文件中的模块名:
import pickle import re with open('model.pt', 'rb') as f: data = f.read() # 将旧模块名替换为当前存在的模块名 new_data = re.sub(b'pytorch_lightning.accelerators.cuda', b'pytorch_lightning.accelerators', data) with open('fixed_model.pt', 'wb') as f: f.write(new_data) # 尝试加载修复后的模型 model = torch.load('fixed_model.pt')
注意:这种方法需谨慎操作,确保替换的字节长度一致,或使用专门的pickle修改工具(如pickle-mixin)避免破坏文件结构。
内容的提问来源于stack exchange,提问作者Janik
相关产品推荐
相关产品推荐

