diff --git a/mmengine/hooks/logger_hook.py b/mmengine/hooks/logger_hook.py index fa0b79dcf9..b7143b9930 100644 --- a/mmengine/hooks/logger_hook.py +++ b/mmengine/hooks/logger_hook.py @@ -262,6 +262,7 @@ def after_val_epoch(self, epoch = 0 else: epoch = runner.epoch + tag['epoch'] = epoch runner.visualizer.add_scalars( tag, step=epoch, file_path=self.json_log_path) else: diff --git a/tests/test_hooks/test_logger_hook.py b/tests/test_hooks/test_logger_hook.py index 52b8bc1fa3..7088aa3262 100644 --- a/tests/test_hooks/test_logger_hook.py +++ b/tests/test_hooks/test_logger_hook.py @@ -120,6 +120,8 @@ def test_after_train_iter(self): def test_after_val_epoch(self): logger_hook = LoggerHook() runner = MagicMock() + runner._train_loop = object() + runner.epoch = 3 # Test when `log_metric_by_epoch` is True runner.log_processor.get_log_after_epoch = MagicMock( return_value=({ @@ -135,7 +137,8 @@ def test_after_val_epoch(self): call({ 'time': 1, 'datatime': 1, - 'acc': 0.8 + 'acc': 0.8, + 'epoch': 3 }, **args), ] self.assertEqual( @@ -157,7 +160,8 @@ def test_after_val_epoch(self): call({ 'time': 1, 'datatime': 1, - 'acc': 0.8 + 'acc': 0.8, + 'epoch': 3 }, **args), call({ 'time': 5,