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

PyTorch Lightning中on_save_checkpoint方法未被调用问题求助

PyTorch Lightning on_save_checkpoint方法未触发问题
  • 项目处于开发阶段,需切换至on_save_checkpointNotWorking分支查看代码。
  • 实现了继承自pytorch lightning LightningModule的BrazingTorch类,路径为brazingTorchFolder/brazingTorch.py;该类的on_save_checkpoint方法定义在父类文件brazingTorchFolder/brazingTorchParents/saveLoad.py中。
  • 执行.fit方法(包含training_step等完整训练流程)时,on_save_checkpoint方法从未被调用。
  • 已在.fit中配置ModelCheckpoint回调,模型能正常保存,但on_save_checkpoint始终不触发。
  • 已确认on_save_checkpoint方法存在于BrazingTorch的继承链中,排查过常规问题后仍未解决。
  • 可通过运行tests\brazingTorchTests\fitTests.py中的.fit方法复现问题,该方法实际调用brazingTorchFolder\brazingTorchParents\modelFitter.py里的.fit方法(与同文件的.baseFit密切相关)。
  • 日志及checkpoint保存路径:tests\brazingTorchTests\NNDummy1\arch1\mainRun_seed71

相关代码如下:

def fit(self, trainDataloader: DataLoader,
      valDataloader: Optional[DataLoader] = None,
      *, lossFuncs: List[nn.modules.loss._Loss],
      seed=None, resume=True, seedSensitive=False,
      addDefaultLogger=True, addDefault_gradientClipping=True,
      preRunTests_force=False, preRunTests_seedSensitive=False,
      preRunTests_lrsToFindBest=None,
      preRunTests_batchSizesToFindBest=None,
      preRunTests_fastDevRunKwargs=None, preRunTests_overfitBatchesKwargs=None,
      preRunTests_profilerKwargs=None, preRunTests_findBestLearningRateKwargs=None,
      preRunTests_findBestBatchSizesKwargs=None,
      **kwargs):

    if not seed:
      seed = self.seed

    self._setLossFuncs_ifNot(lossFuncs)

    architectureName, loggerPath, shouldRun_preRunTests = self._determineShouldRun_preRunTests(
      False, seedSensitive)


    loggerPath = loggerPath.replace('preRunTests', 'mainRun_seed71')

    checkpointCallback = ModelCheckpoint(
      monitor=f"{self._getLossName('val', self.lossFuncs[0])}",
      mode='min', # Save the model when the monitored quantity is minimized
      save_top_k=1, # Save the top model based on the monitored quantity
      every_n_epochs=1, # Checkpoint every 1 epoch
      dirpath=loggerPath, # Directory to save checkpoints
      filename=f'BrazingTorch',
    )
    callbacks_ = [checkpointCallback, StoreEpochData()]
    kwargsApplied = {
      'logger': pl.loggers.TensorBoardLogger(self.modelName, name=architectureName,
                          version='preRunTests'),
      'callbacks': callbacks_, }

    return self.baseFit(trainDataloader=trainDataloader, valDataloader=valDataloader,
              addDefaultLogger=addDefaultLogger,
              addDefault_gradientClipping=addDefault_gradientClipping,
              listOfKwargs=[kwargsApplied], **kwargs)

  @argValidator
  def baseFit(self, trainDataloader: DataLoader,
        valDataloader: Union[DataLoader, None] = None,
        addDefaultLogger=True, addDefault_gradientClipping=True,
        listOfKwargs: List[dict] = None,
        **kwargs):

    # cccUsage
    # - this method accepts kwargs related to trainer, trainer.fit, and self.log and
    # pass them accordingly
    # - the order in listOfKwargs is important
    # - _logOptions phase based values feature:
    #      - args related to self.log may be a dict with these keys 'train', 'val', 'test',
    #          'predict' or 'else'
    #      - this way u can specify what phase use what values and if not specified with
    #        'else' it's gonna know

    # put together all kwargs user wants to pass to trainer, trainer.fit, and self.log
    listOfKwargs = listOfKwargs or []
    listOfKwargs.append(kwargs)
    allUserKwargs = {}
    for kw in listOfKwargs:
      self._plKwargUpdater(allUserKwargs, kw)

    # add default logger if allowed and no logger is passes
    # because by default we are logging some metrics
    if addDefaultLogger and 'logger' not in allUserKwargs:
      allUserKwargs['logger'] = pl.loggers.TensorBoardLogger(self.modelName)
      # bugPotentialCheck1
      # shouldn't this default logger have architectureName

    appliedKwargs = self._getArgsRelated_toEachMethodSeparately(allUserKwargs)

    notAllowedArgs = ['self', 'overfit_batches', 'name', 'value']
    self._removeNotAllowedArgs(allUserKwargs, appliedKwargs, notAllowedArgs)

    self._warnForNotUsedArgs(allUserKwargs, appliedKwargs)

    # add gradient clipping by default
    if not self.noAdditionalOptions and addDefault_gradientClipping \
        and 'gradient_clip_val' not in appliedKwargs['trainer']:
      appliedKwargs['trainer']['gradient_clip_val'] = 0.1
      Warn.info('gradient_clip_val is not provided to fit;' + \
           ' so by default it is set to default "0.1"' + \
           '\nto cancel it, you may either pass noAdditionalOptions=True to model or ' + \
           'pass addDefault_gradientClipping=False to fit method.' + \
           '\nor set another value to "gradient_clip_val" in kwargs passed to fit method.')

    trainer = pl.Trainer(**appliedKwargs['trainer'])

    self._logOptions = appliedKwargs['log']

    if 'train_dataloaders' in appliedKwargs['trainerFit']:
      del appliedKwargs['trainerFit']['train_dataloaders']
    if 'val_dataloaders' in appliedKwargs['trainerFit']:
      del appliedKwargs['trainerFit']['val_dataloaders']
    trainer.fit(self, trainDataloader, valDataloader, **appliedKwargs['trainerFit'])

    self._logOptions = {}
    return trainer
  def on_save_checkpoint(self, checkpoint: dict):
    # reimplement this method to save additional information to the checkpoint

    # Add additional information to the checkpoint
    checkpoint['brazingTorch'] = {
      '_initArgs': self._initArgs,
      'allDefinitions': self.allDefinitions,
      'warnsFrom_getAllNeededDefinitions': self.warnsFrom_getAllNeededDefinitions,
    }
    return checkpoint

内容的提问来源于stack exchange,提问作者Farhang Amaji

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 10:04:51