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

加载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对类路径的依赖,只加载参数字典:

  1. 重新定义与训练时完全一致的PyTorch Lightning模型类(包括层结构、参数初始化等)
  2. 提取并加载模型参数:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 22:15:49