加载BERTSUM训练模型时遇ModuleNotFoundError问题求助
解决BERTSUM模型加载时的ModuleNotFoundError问题
问题原因
你保存的checkpoint中包含了BERTSUM自定义的优化器对象(optim字段),torch.load反序列化时需要找到对应的models.optimizers模块,而当前项目中缺少该模块,因此触发报错。
解决方法
方法1:复制自定义优化器文件到当前项目
从BERTSUM项目的models/目录下,将optimizers.py文件复制到你当前项目的models/目录下(没有则新建该目录),确保目录结构为models/optimizers.py,让Python解释器能找到该模块。
方法2:加载时跳过优化器(仅推理场景)
如果只是加载模型做推理,不需要继续训练,可修改加载代码,临时添加BERTSUM项目路径到Python环境变量,加载后只提取模型参数和配置:
import sys import torch # 替换为你的BERTSUM项目根目录路径 sys.path.append("/path/to/your/BertSum") checkpoint_path = '~/Desktop/fyp/models/bert_transformer/model_step_44000.pt' checkpoint = torch.load(checkpoint_path, map_location=torch.device('cpu')) # 只提取需要的模型参数和配置,丢弃优化器 model_state_dict = checkpoint['model'] args = checkpoint['opt'] # 后续用model_state_dict初始化你的模型即可
方法3:修改保存逻辑(重新保存模型)
如果可以重新训练并保存模型,修改BERTSUM的保存代码,去掉checkpoint中的optim字段,避免序列化优化器对象:
def _save(self, step): real_model = self.model model_state_dict = real_model.state_dict() checkpoint = { 'model': model_state_dict, 'opt': self.args, # 移除'optim': self.optim这一行 } checkpoint_path = os.path.join(self.args.model_path, 'model_step_%d.pt' % step) logger.info("Saving checkpoint %s" % checkpoint_path) if not os.path.exists(checkpoint_path): torch.save(checkpoint, checkpoint_path) return checkpoint, checkpoint_path
内容的提问来源于stack exchange,提问作者ivan
相关产品推荐
相关产品推荐

