diff --git a/mmengine/hooks/checkpoint_hook.py b/mmengine/hooks/checkpoint_hook.py index 92a4867bb9..62194cf138 100644 --- a/mmengine/hooks/checkpoint_hook.py +++ b/mmengine/hooks/checkpoint_hook.py @@ -326,6 +326,11 @@ def before_train(self, runner) -> None: self.keep_ckpt_ids: deque = deque(keep_ckpt_ids, self.max_keep_ckpts) + def before_val(self, runner) -> None: + """Initialize checkpoint state when validation runs standalone.""" + if not hasattr(self, 'file_backend'): + self.before_train(runner) + def after_train_epoch(self, runner) -> None: """Save the checkpoint and synchronize buffers after each epoch. diff --git a/tests/test_hooks/test_checkpoint_hook.py b/tests/test_hooks/test_checkpoint_hook.py index 13914341f7..dc02a76e81 100644 --- a/tests/test_hooks/test_checkpoint_hook.py +++ b/tests/test_hooks/test_checkpoint_hook.py @@ -336,7 +336,6 @@ def test_after_val_epoch(self): petrel_client = MagicMock() for by_epoch, cfg in [(True, self.epoch_based_cfg), (False, self.iter_based_cfg)]: - isfile = MagicMock(return_value=True) self.clear_work_dir() with patch.dict(sys.modules, {'petrel_client': petrel_client}), \ patch('mmengine.fileio.backends.PetrelBackend.put') as put_mock, \ @@ -362,6 +361,17 @@ def test_after_val_epoch(self): isfile.assert_called_once() remove_mock.assert_called_once() + def test_before_val_initializes_save_best(self): + runner = self.build_runner(self.epoch_based_cfg) + checkpoint_hook = CheckpointHook(save_best='acc') + + checkpoint_hook.before_val(runner) + checkpoint_hook.after_val_epoch(runner, {'acc': 0.5}) + + self.assertEqual(runner.message_hub.get_info('best_score'), 0.5) + self.assertTrue( + osp.isfile(osp.join(runner.work_dir, 'best_acc_epoch_0.pth'))) + def test_after_train_epoch(self): cfg = copy.deepcopy(self.epoch_based_cfg) runner = self.build_runner(cfg)