PyTorch Forecasting lr_find触发KeyError: 'radam_buffer'问题
解决PyTorch Forecasting中lr_find报错KeyError: 'radam_buffer'
问题场景
参考PyTorch Forecasting官方文档复现DeepAR模型时,执行以下学习率查找代码持续报错:
# 查找最优学习率 res = trainer.tuner.lr_find( net, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader, min_lr=1e-5, max_lr=1e01, early_stop_threshold=100)
错误信息
File ~\anaconda3\envs\pytorch_forecasting\lib\site-packages\pytorch_lightning\trainer\connectors\checkpoint_connector.py:397, in CheckpointConnector.restore_optimizers(self) 394 return 396 # 恢复优化器状态 --> 397 self.trainer.strategy.load_optimizer_state_dict(self._loaded_checkpoint) File ~\anaconda3\envs\pytorch_forecasting\lib\site-packages\pytorch_lightning\strategies\strategy.py:368, in Strategy.load_optimizer_state_dict(self, checkpoint) 366 optimizer_states = checkpoint["optimizer_states"] 367 for optimizer, opt_state in zip(self.optimizers, optimizer_states): --> 368 optimizer.load_state_dict(opt_state) 369 _optimizer_to_device(optimizer, self.root_device) File ~\anaconda3\envs\pytorch_forecasting\lib\site-packages\torch\optim\optimizer.py:244, in Optimizer.load_state_dict(self, state_dict) 241 return new_group 242 param_groups = [ 243 update_group(g, ng) for g, ng in zip(groups, saved_groups)] --> 244 self.__setstate__({'state': state, 'param_groups': param_groups}) File ~\anaconda3\envs\pytorch_forecasting\lib\site-packages\pytorch_forecasting\optim.py:133, in Ranger.__setstate__(self, state) 131 def __setstate__(self, state: dict) -> None: 132 super().__setstate__(state) --> 133 self.radam_buffer = state["radam_buffer"] 134 self.alpha = state["alpha"] 135 self.k = state["k"] KeyError: 'radam_buffer'
解决方案
1. 更换为PyTorch原生优化器
在初始化DeepAR模型时,显式指定使用Adam等原生优化器,规避Ranger优化器的兼容性问题:
from pytorch_forecasting import DeepAR net = DeepAR( # 其他模型参数... optimizer="Adam", optimizer_kwargs={"lr": 1e-3} )
2. 修改Ranger优化器的状态加载逻辑
找到环境中pytorch_forecasting/optim.py文件里的Ranger类,修改其__setstate__方法,添加键不存在时的默认值处理:
def __setstate__(self, state: dict) -> None: super().__setstate__(state) # 使用get方法获取值,不存在时用默认值兜底 self.radam_buffer = state.get("radam_buffer", [[None, None, None] for _ in range(10)]) self.alpha = state.get("alpha", 0.5) self.k = state.get("k", 6)
3. 更新PyTorch Forecasting到最新版本
执行以下命令升级库,新版本大概率已修复该兼容问题:
pip install --upgrade pytorch-forecasting
内容的提问来源于stack exchange,提问作者Diego Alvarez
相关产品推荐
相关产品推荐

