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
相关产品推荐
相关产品推荐

