diff --git a/mmengine/model/utils.py b/mmengine/model/utils.py index c78ea3134d..59a4f47841 100644 --- a/mmengine/model/utils.py +++ b/mmengine/model/utils.py @@ -248,6 +248,7 @@ def convert_sync_batchnorm(module: nn.Module, module_output.running_mean = module.running_mean module_output.running_var = module.running_var module_output.num_batches_tracked = module.num_batches_tracked + module_output.training = module.training if hasattr(module, 'qconfig'): module_output.qconfig = module.qconfig for name, child in module.named_children(): diff --git a/tests/test_hooks/test_checkpoint_hook.py b/tests/test_hooks/test_checkpoint_hook.py index 13914341f7..b078a25fcb 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, \ diff --git a/tests/test_model/test_convert_sync_batchnorm_training.py b/tests/test_model/test_convert_sync_batchnorm_training.py new file mode 100644 index 0000000000..ef7fbf16c0 --- /dev/null +++ b/tests/test_model/test_convert_sync_batchnorm_training.py @@ -0,0 +1,14 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import torch.nn as nn + +from mmengine.model import convert_sync_batchnorm + + +def test_convert_sync_batchnorm_keeps_training_state(): + bn = nn.BatchNorm2d(4) + bn.eval() + + sync_bn = convert_sync_batchnorm(bn) + + assert isinstance(sync_bn, nn.SyncBatchNorm) + assert not sync_bn.training