仅用torch.save(model)保存T5模型后如何加载及恢复?
解决T5模型加载问题的方法
问题原因
- 用
torch.save(model, 'trained_model')保存整个模型对象时,会序列化模型类的路径信息,若当前环境的transformers版本与训练时不一致,或模块结构变更(如旧版transformers.modeling_t5路径在新版中改为transformers.models.t5.modeling_t5),就会触发No module named 'transformers.modeling_t5'错误。 T5ForConditionalGeneration.from_pretrained()要求加载的是包含config.json、pytorch_model.bin等文件的标准目录,而非单个模型文件,因此会提示配置文件无效。
无需重新训练的恢复方法
方法一:直接加载权重到空模型
- 安装与训练时完全相同版本的transformers:
# 先查看训练环境的transformers版本 pip show transformers # 在当前环境安装对应版本 pip install transformers==x.x.x - 初始化与训练时结构一致的空模型,再加载保存的权重:
from transformers import T5ForConditionalGeneration import torch # 初始化原始结构的空模型 raw_model = 'rut5-base-multitask' model = T5ForConditionalGeneration.from_pretrained(raw_model) # 加载单个模型文件的状态字典 saved_model = torch.load('trained_model', map_location='cpu') # 可指定cuda:0等设备 # 将权重加载到空模型 model.load_state_dict(saved_model.state_dict()) # 验证模型 model.eval()
方法二:转换为transformers标准保存格式
如果方法一仍无法加载,可先修复模块依赖后,将单个文件转换为标准格式:
- 加载模型(需先通过安装对应版本transformers解决模块路径问题):
import torch from transformers import T5ForConditionalGeneration model = torch.load('trained_model', map_location='cpu') - 保存为transformers标准格式:
# 保存到指定目录,自动生成config.json、pytorch_model.bin等文件 model.save_pretrained('./trained_model_standard') - 后续即可用标准方式加载:
model = T5ForConditionalGeneration.from_pretrained('./trained_model_standard')
内容的提问来源于stack exchange,提问作者Dion
相关产品推荐
相关产品推荐

