基于PyTorch Lightning的TFT模型Checkpoint重训练无进展问题
背景与问题
我基于PyTorch Lightning的pytorch-forecasting库训练TFT模型进行时序预测,训练流程分为三个阶段:
- 用Optuna进行超参数优化(HPO)
- 执行训练+验证
- 用全数据重训练(无验证)
为了缩短重训练耗时,尝试复用阶段2的Checkpoint做增量训练,减少重训练的epoch数。但调用自定义的fit_model()方法进行重训练时,日志显示已成功加载Checkpoint,却未执行后续训练,返回的best_model_path仍是阶段2的旧路径。使用的库版本为pytorch-lightning==1.6.5、pytorch-forecasting==0.9.0。
核心代码
def fit_model(self, **kwargs): ... to_retrain = kwargs.get('to_retrain', False) ckpt_path = kwargs.get('ckpt_path', None) trainer = self._get_trainer(cluster_id, gpu_id, to_retrain) # 返回pl.Trainer实例 tft_lightning_module = self._prepare_for_training(cluster_id, to_retrain) train_dtloaders = ... val_dtloaders = ... if not to_retrain: trainer.fit( tft_lightning_module, train_dataloaders=train_dtloaders, val_dataloaders=val_dtloaders ) else: trainer.fit( tft_lightning_module, train_dataloaders=train_dtloaders, val_dataloaders=val_dtloaders, ckpt_path=ckpt_path ) best_model_path = trainer.checkpoint_callback.best_model_path return best_model_path
加载Checkpoint的日志
Restored all states from the checkpoint file at /tft/incremental_training/tft_training_20230206/171049/lightning_logs_3/lightning_logs/version_0/checkpoints/epoch=4-step=5.ckpt
解决方案
1. 调整Trainer的max_epochs参数
重训练时,_get_trainer()返回的Trainer实例max_epochs可能与Checkpoint中记录的已训练epoch数相等,导致模型判定训练已完成,直接终止。
- 修复方式:根据
to_retrain参数设置更大的max_epochs,确保重训练有足够的epoch可以执行:
def _get_trainer(self, cluster_id, gpu_id, to_retrain): # 原有其他配置逻辑 if to_retrain: # 假设阶段2训练了5个epoch,重训练设置为10个epoch max_epochs = 10 else: max_epochs = 5 trainer = pl.Trainer( max_epochs=max_epochs, # 其他参数(如accelerator、callbacks等) ) return trainer
2. 单独配置重训练的CheckpointCallback
阶段2的CheckpointCallback可能与重训练的保存目录冲突,导致新的Checkpoint未被保存,best_model_path仍指向旧路径。
- 修复方式:重训练时单独配置CheckpointCallback,指定独立的保存目录和命名规则:
from pytorch_lightning.callbacks import ModelCheckpoint def _get_trainer(self, cluster_id, gpu_id, to_retrain): if to_retrain: # 重训练的Checkpoint配置 checkpoint_callback = ModelCheckpoint( dirpath=f"./tft_retrain_ckpts/{cluster_id}", filename="retrain-{epoch:02d}-{step:02d}", save_top_k=1, # 重训练无验证时,监控训练损失 monitor="train_loss" ) trainer = pl.Trainer( max_epochs=10, callbacks=[checkpoint_callback], check_val_every_n_epoch=0, # 关闭验证步骤 # 其他参数 ) else: # 阶段2的Checkpoint配置 checkpoint_callback = ModelCheckpoint( dirpath=f"./tft_val_ckpts/{cluster_id}", filename="val-{epoch:02d}-{val_loss:.2f}", save_top_k=1, monitor="val_loss" ) trainer = pl.Trainer( max_epochs=5, callbacks=[checkpoint_callback], # 其他参数 ) return trainer
3. 重训练时移除验证数据加载器
用户提到重训练是无验证的,但代码中仍传入val_dtloaders,可能导致模型因验证指标未提升而不保存新Checkpoint,甚至触发提前停止逻辑。
- 修复方式:重训练时不再传入
val_dataloaders:
# fit_model方法的else分支中 trainer.fit( tft_lightning_module, train_dataloaders=train_dtloaders, # 移除val_dataloaders参数 ckpt_path=ckpt_path )
4. 手动加载Checkpoint并设置起始epoch
如果上述方法仍无效,可以手动加载Checkpoint到模型,显式设置起始epoch,避免Trainer误判训练状态:
else: tft_lightning_module = self._prepare_for_training(cluster_id, to_retrain) # 手动加载Checkpoint状态 tft_lightning_module = tft_lightning_module.load_from_checkpoint(ckpt_path) # 设置起始epoch(比如阶段2训练到了epoch4,重训练从epoch5开始) tft_lightning_module.current_epoch = 4 # 不需要再传入ckpt_path参数 trainer.fit( tft_lightning_module, train_dataloaders=train_dtloaders )
内容的提问来源于stack exchange,提问作者Bitswazsky

