You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于PyTorch Lightning的TFT模型Checkpoint重训练无进展问题

问题:PyTorch Lightning增量训练TFT模型未执行训练且返回旧Checkpoint路径

背景与问题

我基于PyTorch Lightning的pytorch-forecasting库训练TFT模型进行时序预测,训练流程分为三个阶段:

  1. 用Optuna进行超参数优化(HPO)
  2. 执行训练+验证
  3. 用全数据重训练(无验证)

为了缩短重训练耗时,尝试复用阶段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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 06:45:30