From a569ff078385d496e67488748c3fa18f4a446e8e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=8B=E5=8D=97?= Date: Mon, 24 Jan 2022 17:04:17 +0800 Subject: [PATCH] bugfix: Model amp args differ from amp.build_train_network --- mindspore/python/mindspore/train/model.py | 28 +++++++++-------------- 1 file changed, 11 insertions(+), 17 deletions(-) diff --git a/mindspore/python/mindspore/train/model.py b/mindspore/python/mindspore/train/model.py index 056b1222c32..3f1e0c9a879 100644 --- a/mindspore/python/mindspore/train/model.py +++ b/mindspore/python/mindspore/train/model.py @@ -180,7 +180,7 @@ class Model: self._optimizer = optimizer self._loss_scale_manager = None self._loss_scale_manager_set = False - self._keep_bn_fp32 = True + self._keep_bn_fp32 = None self._check_kwargs(kwargs) self._amp_level = amp_level self._boost_level = boost_level @@ -215,8 +215,6 @@ class Model: raise ValueError("For 'Model', the '**kwargs' argument should be empty when network is a GraphCell.") def _process_amp_args(self, kwargs): - if self._amp_level in ["O0", "O3"]: - self._keep_bn_fp32 = False if 'keep_batchnorm_fp32' in kwargs: self._keep_bn_fp32 = kwargs['keep_batchnorm_fp32'] if 'loss_scale_manager' in kwargs: @@ -271,21 +269,17 @@ class Model: raise ValueError("The argument 'optimizer' can not be None when set 'loss_scale_manager'.") if self._optimizer: + amp_config = {} if self._loss_scale_manager_set: - network = amp.build_train_network(network, - self._optimizer, - self._loss_fn, - level=self._amp_level, - boost_level=self._boost_level, - loss_scale_manager=self._loss_scale_manager, - keep_batchnorm_fp32=self._keep_bn_fp32) - else: - network = amp.build_train_network(network, - self._optimizer, - self._loss_fn, - level=self._amp_level, - boost_level=self._boost_level, - keep_batchnorm_fp32=self._keep_bn_fp32) + amp_config['loss_scale_manager'] = self._loss_scale_manager + if self._keep_bn_fp32 is not None: + amp_config['keep_batchnorm_fp32'] = self._keep_bn_fp32 + network = amp.build_train_network(network, + self._optimizer, + self._loss_fn, + level=self._amp_level, + boost_level=self._boost_level, + **amp_config) elif self._loss_fn: if self._parallel_mode in (ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL): network = _VirtualDatasetCell(network)