diff --git a/mmengine/runner/loops.py b/mmengine/runner/loops.py index 5a678db7b9..dc6cf85ff0 100644 --- a/mmengine/runner/loops.py +++ b/mmengine/runner/loops.py @@ -26,7 +26,9 @@ class EpochBasedTrainLoop(BaseLoop): dataloader (Dataloader or dict): A dataloader object or a dict to build a dataloader. max_epochs (int): Total training epochs. - val_begin (int): The epoch that begins validating. + val_begin (int): The epoch that begins validating. If it is set to + 0, an additional validation is performed before the first epoch + is trained, which allows the model to be evaluated at epoch 0. Defaults to 1. val_interval (int): Validation interval. Defaults to 1. dynamic_intervals (List[Tuple[int, int]], optional): The @@ -94,6 +96,12 @@ def run(self) -> torch.nn.Module: """Launch training.""" self.runner.call_hook('before_train') + # `val_begin` is 1-based for the validation performed after an epoch. + # Allowing it to be 0 additionally validates at epoch 0, i.e. before + # any training step has been taken. + if self.val_begin <= 0 and self.runner.val_loop is not None: + self.runner.val_loop.run() + while self._epoch < self._max_epochs and not self.stop_training: self.run_epoch() diff --git a/tests/test_runner/test_runner.py b/tests/test_runner/test_runner.py index 7e105f0895..02060399fb 100644 --- a/tests/test_runner/test_runner.py +++ b/tests/test_runner/test_runner.py @@ -1835,6 +1835,57 @@ def train_step(self, *args, **kwargs): runner.train() self.assertEqual(runner.iter, 3 * 2) + def test_val_begin_zero(self): + # `val_begin=0` should additionally validate before the first epoch is + # trained, so that the model can be evaluated at epoch 0. See #1448. + val_epochs = [] + val_iters = [] + + @HOOKS.register_module(force=True) + class TestValBeginZeroHook(Hook): + + def before_val_epoch(self, runner): + val_epochs.append(runner.epoch) + val_iters.append(runner.iter) + + cfg = copy.deepcopy(self.epoch_based_cfg) + cfg.experiment_name = 'test_val_begin_zero' + cfg.custom_hooks = [dict(type='TestValBeginZeroHook', priority=50)] + cfg.train_cfg = dict( + by_epoch=True, max_epochs=2, val_interval=1, val_begin=0) + runner = Runner.from_cfg(cfg) + runner.train() + + # epoch 0 is validated before any training step has been taken, and + # the regular per-epoch validation is unaffected + self.assertEqual(val_epochs, [0, 1, 2]) + self.assertEqual(val_iters, [0, 4, 8]) + HOOKS.module_dict.pop('TestValBeginZeroHook') + + # the default behaviour must stay unchanged: validation only happens + # after an epoch has been trained + val_epochs = [] + val_iters = [] + + @HOOKS.register_module(force=True) + class TestValBeginDefaultHook(Hook): + + def before_val_epoch(self, runner): + val_epochs.append(runner.epoch) + val_iters.append(runner.iter) + + cfg = copy.deepcopy(self.epoch_based_cfg) + cfg.experiment_name = 'test_val_begin_default' + cfg.custom_hooks = [dict(type='TestValBeginDefaultHook', priority=50)] + cfg.train_cfg = dict( + by_epoch=True, max_epochs=2, val_interval=1, val_begin=1) + runner = Runner.from_cfg(cfg) + runner.train() + + self.assertEqual(val_epochs, [1, 2]) + self.assertEqual(val_iters, [4, 8]) + HOOKS.module_dict.pop('TestValBeginDefaultHook') + @skipIf( SKIP_TEST_COMPILE, reason='torch.compile is not valid, please install PyTorch>=2.0.0')