如何恢复与PyTorch Lightning集成的Ray Tune生成的最佳检查点?
问题原因
PyTorch Lightning的load_from_checkpoint方法默认仅加载模型权重,不会自动获取你训练阶段传入构造函数的config参数。你调用方法时没有主动传入config,就会触发默认值None,访问config["lr"]时报下标访问错误。
解决方法
方法1:直接传入Ray Tune返回的最优配置(无需修改模型代码)
Ray Tune的analysis对象自带best_config属性,存储了最优试验对应的超参数配置,你可以直接作为关键字参数传给load_from_checkpoint,方法会自动把额外参数传递给模型构造函数:
MyLightningModel.load_from_checkpoint( os.path.join(analysis.best_checkpoint, "checkpoint"), config=analysis.best_config )
该方案适合快速验证,不需要调整原有训练逻辑。
方法2:在模型中保存超参数(更规范的长期方案)
你可以在模型的__init__方法中调用PyTorch Lightning提供的save_hyperparameters()方法,它会自动把构造函数的参数存入检查点,加载时不需要额外传参就能自动恢复:
修改后的模型代码示例:
class MyLightningModel (pl.LightningModule): def __init__(self, config=None): # 将构造函数参数存入检查点 self.save_hyperparameters() self.lr = config["lr"] self.batch_size = config["batch_size"] self.layer_size = config["layer_size"] super(MyLightningModel , self).__init__() self.lstm = nn.LSTM(768, self.layer_size, num_layers=1, bidirectional=False) self.out = nn.Linear(self.layer_size, 1)
修改后训练生成的检查点,后续加载时可以直接用你原有代码,不需要额外传递config参数。
内容的提问来源于stack exchange,提问作者Luca Guarro
相关产品推荐
相关产品推荐

