diff --git a/mindspore/ccsrc/frontend/optimizer/irpass.cc b/mindspore/ccsrc/frontend/optimizer/irpass.cc index 09d3dba818c..d3ff88e572f 100644 --- a/mindspore/ccsrc/frontend/optimizer/irpass.cc +++ b/mindspore/ccsrc/frontend/optimizer/irpass.cc @@ -177,7 +177,8 @@ OptimizeIRPassLib::OptimizeIRPassLib() { // Accelerated Algorithm less_batch_normalization_ = - MakeSubstitution(std::make_shared(), "less_batch_normalization", prim::kPrimAdd); + MakeSubstitution(std::make_shared(), "less_batch_normalization", + {prim::kPrimAdd, prim::kPrimRelu6, prim::kPrimMatMul, prim::kPrimMakeTuple, prim::kPrimMaxPool}); // inline inline_ = MakeSubstitution(std::make_shared(), "inline", IsCNodeGraph); diff --git a/mindspore/ccsrc/frontend/optimizer/irpass/less_batch_normalization.cc b/mindspore/ccsrc/frontend/optimizer/irpass/less_batch_normalization.cc index 9ce30eeaa7f..77ca4b7f752 100644 --- a/mindspore/ccsrc/frontend/optimizer/irpass/less_batch_normalization.cc +++ b/mindspore/ccsrc/frontend/optimizer/irpass/less_batch_normalization.cc @@ -31,8 +31,8 @@ constexpr auto kFirstBranchPattern1 = 12; constexpr auto kSecondBranchPattern1 = 3; constexpr auto kFirstBranchStartIndexPattern1 = 4; constexpr auto kFirstBranchEndIndexPattern1 = 11; -constexpr auto kSecondBranchStartIndexPattern1 = 12; -constexpr auto kSecondBranchEndIndexPattern1 = 14; +constexpr auto kSecondBranchStartIndexPattern1 = kFirstBranchPattern1; +constexpr auto kSecondBranchEndIndexPattern1 = 2 + kFirstBranchPattern1; const std::vector ResidualStructureBasePattern{ {kFirstBranchPattern1, {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu}, @@ -47,8 +47,8 @@ constexpr auto kFirstBranchPattern2 = 12; constexpr auto kSecondBranchPattern2 = 1; constexpr auto kFirstBranchStartIndexPattern2 = 4; constexpr auto kFirstBranchEndIndexPattern2 = 11; -constexpr auto kSecondBranchStartIndexPattern2 = 12; -constexpr auto kSecondBranchEndIndexPattern2 = 13; +constexpr auto kSecondBranchStartIndexPattern2 = kFirstBranchPattern2; +constexpr auto kSecondBranchEndIndexPattern2 = 1 + kSecondBranchPattern2; const std::vector ResidualStructureShortCutPattern{ {kFirstBranchPattern2, {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu}, @@ -61,8 +61,8 @@ constexpr auto kFirstBranchPattern3 = 11; constexpr auto kSecondBranchPattern3 = 3; constexpr auto kFirstBranchStartIndexPattern3 = 4; constexpr auto kFirstBranchEndIndexPattern3 = 10; -constexpr auto kSecondBranchStartIndexPattern3 = 11; -constexpr auto kSecondBranchEndIndexPattern3 = 13; +constexpr auto kSecondBranchStartIndexPattern3 = kFirstBranchPattern3; +constexpr auto kSecondBranchEndIndexPattern3 = 2 + kFirstBranchPattern3; const std::vector ResidualStructureFirstStepPattern{ {kFirstBranchPattern3, {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, @@ -73,15 +73,13 @@ const std::vector ResidualStructureFirstStepPattern{ {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D}, {kSecondBranchStartIndexPattern3, kSecondBranchEndIndexPattern3}}}; // Pattern 4 -// Add -> BatchNorm -> Conv2D -> Relu ... -> End -// ↘ BatchNorm -> Conv2D -> -> -> -> ↗ constexpr auto kFirstBranchPattern4 = 8; constexpr auto kSecondBranchPattern4 = 3; constexpr auto kFirstBranchStartIndexPattern4 = 4; constexpr auto kFirstBranchEndIndexPattern4 = 6; -constexpr auto kSecondBranchStartIndexPattern4 = 8; -constexpr auto kSecondBranchEndIndexPattern4 = 11; -const std::vector BasicStructureBasePattern{ +constexpr auto kSecondBranchStartIndexPattern4 = kFirstBranchPattern4; +constexpr auto kSecondBranchEndIndexPattern4 = 3 + kFirstBranchPattern4; +const std::vector BasicStructBasePattern{ {kFirstBranchPattern4, {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu}, {kFirstBranchStartIndexPattern4, kFirstBranchEndIndexPattern4}}, @@ -89,37 +87,163 @@ const std::vector BasicStructureBasePattern{ {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D}, {kSecondBranchStartIndexPattern4, kSecondBranchEndIndexPattern4}}}; // Pattern 5 -// Add -> BatchNorm -> Conv2D -> Relu ... -> End -// ↘ -> -> -> -> Relu -> -> -> -> ↗ -constexpr auto kFirstBranchPattern5 = 8; +constexpr auto kFirstBranchPattern5 = 7; constexpr auto kSecondBranchPattern5 = 1; constexpr auto kFirstBranchStartIndexPattern5 = 4; constexpr auto kFirstBranchEndIndexPattern5 = 6; -constexpr auto kSecondBranchStartIndexPattern5 = 8; -constexpr auto kSecondBranchEndIndexPattern5 = 11; -const std::vector BasicStructureShortCutPattern{ +constexpr auto kSecondBranchStartIndexPattern5 = kFirstBranchPattern5; +constexpr auto kSecondBranchEndIndexPattern5 = 3 + kFirstBranchPattern5; +const std::vector BasicStructFirstStepPattern{ {kFirstBranchPattern5, - {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu}, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, + prim::kPrimBatchNorm, prim::kPrimConv2D}, {kFirstBranchStartIndexPattern5, kFirstBranchEndIndexPattern5}}, - {kSecondBranchPattern5, {prim::kPrimRelu}, {kSecondBranchStartIndexPattern5, kSecondBranchEndIndexPattern5}}}; + {kSecondBranchPattern5, {prim::kPrimMaxPool}, {kSecondBranchStartIndexPattern5, kSecondBranchEndIndexPattern5}}}; // Pattern 6 -// Add -> BatchNorm -> Conv2D -> Relu ... -> End -// ↘ -> -> -> -> MaxPool -> -> -> ↗ -constexpr auto kFirstBranchPattern6 = 7; +constexpr auto kFirstBranchPattern6 = 8; constexpr auto kSecondBranchPattern6 = 1; constexpr auto kFirstBranchStartIndexPattern6 = 4; constexpr auto kFirstBranchEndIndexPattern6 = 6; -constexpr auto kSecondBranchStartIndexPattern6 = 7; -constexpr auto kSecondBranchEndIndexPattern6 = 10; -const std::vector BasicStructureFirstStepPattern{ +constexpr auto kSecondBranchStartIndexPattern6 = kFirstBranchPattern6; +constexpr auto kSecondBranchEndIndexPattern6 = 3 + kFirstBranchPattern6; +const std::vector BasicStructShortCutPattern{ {kFirstBranchPattern6, - {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, - prim::kPrimBatchNorm, prim::kPrimConv2D}, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu}, {kFirstBranchStartIndexPattern6, kFirstBranchEndIndexPattern6}}, - {kSecondBranchPattern6, {prim::kPrimMaxPool}, {kSecondBranchStartIndexPattern6, kSecondBranchEndIndexPattern6}}}; -static const std::vector> kNeedMatchPattern = { - ResidualStructureBasePattern, ResidualStructureShortCutPattern, ResidualStructureFirstStepPattern, - BasicStructureBasePattern, BasicStructureShortCutPattern, BasicStructureFirstStepPattern}; + {kSecondBranchPattern6, {prim::kPrimRelu}, {kSecondBranchStartIndexPattern6, kSecondBranchEndIndexPattern6}}}; +// Pattern 7 +constexpr auto kFirstBranchPattern7 = 1; +constexpr auto kSecondBranchPattern7 = 13; +constexpr auto kFirstBranchStartIndexPattern7 = SIZE_MAX; +constexpr auto kFirstBranchEndIndexPattern7 = SIZE_MAX; +constexpr auto kSecondBranchStartIndexPattern7 = 7; +constexpr auto kSecondBranchEndIndexPattern7 = 10; +const std::vector InvertedResidualShortCutPattern{ + {kFirstBranchPattern7, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm}, + {kFirstBranchStartIndexPattern7, kFirstBranchEndIndexPattern7}}, + {kSecondBranchPattern7, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, + prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, + prim::kPrimConv2D, prim::kPrimTupleGetItem, prim::kPrimBatchNorm}, + {kSecondBranchStartIndexPattern7, kSecondBranchEndIndexPattern7}}}; +// Pattern 8 +constexpr auto kFirstBranchPattern8 = 4; +constexpr auto kFirstBranchStartIndexPattern8 = 0; +constexpr auto kFirstBranchEndIndexPattern8 = 3; +const std::vector InvertedResidualPattern{ + {kFirstBranchPattern8, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimAdd}, + {kFirstBranchStartIndexPattern8, kFirstBranchEndIndexPattern8}}}; +// Pattern 9 +constexpr auto kFirstBranchPattern9 = 1; +constexpr auto kSecondBranchPattern9 = 12; +constexpr auto kFirstBranchStartIndexPattern9 = SIZE_MAX; +constexpr auto kFirstBranchEndIndexPattern9 = SIZE_MAX; +constexpr auto kSecondBranchStartIndexPattern9 = 7; +constexpr auto kSecondBranchEndIndexPattern9 = 10; +const std::vector InvertedResidualShortCutPattern2{ + {kFirstBranchPattern9, {prim::kPrimAdd}, {kFirstBranchStartIndexPattern9, kFirstBranchEndIndexPattern9}}, + {kSecondBranchPattern9, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, + prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, + prim::kPrimConv2D, prim::kPrimAdd}, + {kSecondBranchStartIndexPattern9, kSecondBranchEndIndexPattern9}}}; +// Pattern 10 +constexpr auto kFirstBranchPattern10 = 5; +constexpr auto kFirstBranchStartIndexPattern10 = 0; +constexpr auto kFirstBranchEndIndexPattern10 = 4; +const std::vector InvertedResidualPattern2{ + {kFirstBranchPattern10, + {prim::kPrimReduceMean, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D}, + {kFirstBranchStartIndexPattern10, kFirstBranchEndIndexPattern10}}}; +// Pattern 11 +constexpr auto kFirstBranchPattern11 = 17; +constexpr auto kFirstBranchStartIndexPattern11 = 3; +constexpr auto kFirstBranchEndIndexPattern11 = 6; +const std::vector InvertedResidualPattern3{ + {kFirstBranchPattern11, + {prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, + prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, + prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, + prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D}, + {kFirstBranchStartIndexPattern11, kFirstBranchEndIndexPattern11}}}; +// Pattern 12 +constexpr auto kFirstBranchPattern12 = 1; +constexpr auto kSecondBranchPattern12 = 9; +constexpr auto kFirstBranchStartIndexPattern12 = SIZE_MAX; +constexpr auto kFirstBranchEndIndexPattern12 = SIZE_MAX; +constexpr auto kSecondBranchStartIndexPattern12 = kFirstBranchPattern12 + 5; +constexpr auto kSecondBranchEndIndexPattern12 = kFirstBranchPattern12 + 8; +const std::vector DenseBlockShortCutPattern{ + {kFirstBranchPattern12, {prim::kPrimConcat}, {kFirstBranchStartIndexPattern12, kFirstBranchEndIndexPattern12}}, + {kSecondBranchPattern12, + {prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, + prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat}, + {kSecondBranchStartIndexPattern12, kSecondBranchEndIndexPattern12}}}; +// Pattern 13 +constexpr auto kFirstBranchPattern13 = 5; +constexpr auto kFirstBranchStartIndexPattern13 = 0; +constexpr auto kFirstBranchEndIndexPattern13 = 4; +const std::vector DenseBlockPattern{ + {kFirstBranchPattern13, + {prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat}, + {kFirstBranchStartIndexPattern13, kFirstBranchEndIndexPattern13}}}; +// Pattern 14 +constexpr auto kFirstBranchPattern14 = 9; +constexpr auto kSecondBranchPattern14 = 1; +constexpr auto kFirstBranchStartIndexPattern14 = 5; +constexpr auto kFirstBranchEndIndexPattern14 = 8; +constexpr auto kSecondBranchStartIndexPattern14 = SIZE_MAX; +constexpr auto kSecondBranchEndIndexPattern14 = SIZE_MAX; +const std::vector DenseBlockShortCutPattern2{ + {kFirstBranchPattern14, + {prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, + prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat}, + {kFirstBranchStartIndexPattern14, kFirstBranchEndIndexPattern14}}, + {kSecondBranchPattern14, {prim::kPrimConcat}, {kSecondBranchStartIndexPattern14, kSecondBranchEndIndexPattern14}}}; +// Pattern 15 +constexpr auto kFirstBranchPattern15 = 9; +constexpr auto kSecondBranchPattern15 = 1; +constexpr auto kFirstBranchStartIndexPattern15 = 0; +constexpr auto kFirstBranchEndIndexPattern15 = 4; +constexpr auto kSecondBranchStartIndexPattern15 = SIZE_MAX; +constexpr auto kSecondBranchEndIndexPattern15 = SIZE_MAX; +const std::vector DenseBlockPoolPattern{ + {kFirstBranchPattern15, + {prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, + prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimMaxPool}, + {kFirstBranchStartIndexPattern15, kFirstBranchEndIndexPattern15}}, + {kSecondBranchPattern15, {prim::kPrimConcat}, {kSecondBranchStartIndexPattern15, kSecondBranchEndIndexPattern15}}}; +// Pattern 16 +constexpr auto kFirstBranchPattern16 = 1; +constexpr auto kSecondBranchPattern16 = 9; +constexpr auto kFirstBranchStartIndexPattern16 = SIZE_MAX; +constexpr auto kFirstBranchEndIndexPattern16 = SIZE_MAX; +constexpr auto kSecondBranchStartIndexPattern16 = kFirstBranchPattern16; +constexpr auto kSecondBranchEndIndexPattern16 = kFirstBranchPattern16 + 4; +const std::vector DenseBlockPoolPatter2{ + {kFirstBranchPattern16, {prim::kPrimConcat}, {kFirstBranchStartIndexPattern16, kFirstBranchEndIndexPattern16}}, + {kSecondBranchPattern16, + {prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, + prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimMaxPool}, + {kSecondBranchStartIndexPattern16, kSecondBranchEndIndexPattern16}}}; +static const std::vector> kNeedMatchPattern = {ResidualStructureBasePattern, + ResidualStructureShortCutPattern, + ResidualStructureFirstStepPattern, + BasicStructBasePattern, + BasicStructFirstStepPattern, + BasicStructShortCutPattern, + InvertedResidualShortCutPattern, + InvertedResidualPattern, + InvertedResidualShortCutPattern2, + InvertedResidualPattern2, + InvertedResidualPattern3, + DenseBlockShortCutPattern, + DenseBlockPattern, + DenseBlockShortCutPattern2, + DenseBlockPoolPattern, + DenseBlockPoolPatter2}; const std::set kNeedRemoveNodeSet{ prim::kPrimLoad, prim::kPrimRefToEmbed, prim::kPrimApplyMomentum, prim::kPrimMomentum, prim::kPrimApplyFtrl, prim::kPrimSGD, prim::kPrimApplyRMSProp, prim::kPrimAdam}; @@ -286,7 +410,13 @@ AnfNodePtr LessBatchNormalization::operator()(const OptimizerPtr &optimizer, con sum_match_node += std::get<0>(t); total_match_node_.emplace_back(sum_match_node); }); - AnfVisitor::Match(prim::kPrimAdd, {IsCNode, IsCNode})(node); + auto cnode = node->cast(); + if (cnode == nullptr || cnode->inputs().empty()) { + return nullptr; + } + auto prim = GetValueNode(cnode->input(0)); + std::vector funcs(cnode->inputs().size() - 1, IsCNode); + AnfVisitor::Match(prim, funcs)(node); if (is_match_) { break; } diff --git a/mindspore/nn/acc/acc.py b/mindspore/nn/acc/acc.py index 4c8d84a458c..10c10368675 100644 --- a/mindspore/nn/acc/acc.py +++ b/mindspore/nn/acc/acc.py @@ -40,6 +40,7 @@ class AutoAcc: def __init__(self, level, kwargs): if level not in _acc_config_level.keys(): level = 'O0' + self.level = level acc_config = _acc_config_level[level] self._acc_config = acc_config self._fn_flag = True @@ -66,12 +67,11 @@ class AutoAcc: optimizer_process = OptimizerProcess(optimizer) group_params = self._param_processer.assign_parameter_group(network.trainable_params(), self._gradient_groups) - optimizer_process.origin_params = self._param_processer.generate_group_params(group_params) + optimizer_process.origin_params = \ + self._param_processer.generate_group_params(group_params, optimizer_process.origin_params) if self._gc_flag: - parameters = optimizer_process.add_grad_centralization() - else: - parameters = optimizer_process.origin_params - optimizer = optimizer_process.generate_new_optimizer(parameters) + optimizer_process.add_grad_centralization() + optimizer = optimizer_process.generate_new_optimizer() if self._acc_config["grad_freeze"]: freeze_processer = GradientFreeze(self._param_groups, self._freeze_type, diff --git a/mindspore/nn/acc/base.py b/mindspore/nn/acc/base.py index e3d22604415..2cf60cc4dd0 100644 --- a/mindspore/nn/acc/base.py +++ b/mindspore/nn/acc/base.py @@ -13,6 +13,7 @@ # limitations under the License. # ============================================================================ """base process""" +import copy from mindspore.nn.cell import Cell from mindspore.nn.optim import LARS from mindspore import log as logger @@ -58,24 +59,37 @@ class OptimizerProcess: if isinstance(parameters[0], Parameter): logger.warning("Only group parameters support gradient centralization.") - return parameters + return - change_dict = parameters[0] - if 'order_params' in change_dict.keys(): - logger.warning("Only support normal parameters for gradient centralization.") - return parameters + group_params = [] + for group_param in parameters: + if 'order_params' in group_param.keys(): + group_params.append(group_param) + continue + params_gc_value = [] + params_value = [] + for param in group_param['params']: + if 'beta' not in param.name and 'gamma' not in param.name and 'bias' not in param.name: + params_gc_value.append(param) + else: + params_value.append(param) + if params_gc_value: + new_group_param = copy.deepcopy(group_param) + new_group_param['params'] = params_gc_value + new_group_param['grad_centralization'] = True + group_params.append(new_group_param) + if params_value: + new_group_param = copy.deepcopy(group_param) + new_group_param['params'] = params_value + group_params.append(new_group_param) + self.origin_params = group_params - change_dict['grad_centralization'] = True - self.origin_params[0] = change_dict - - return self.origin_params - - def generate_new_optimizer(self, params): + def generate_new_optimizer(self): """Generate new optimizer.""" if not self.is_lars: - opt = self.opt_class(params=params, **self.opt_init_args) + opt = self.opt_class(params=self.origin_params, **self.opt_init_args) else: - opt = LARS(self.opt_class(params=params, **self.opt_init_args), **self.lars_init_args) + opt = LARS(self.opt_class(params=self.origin_params, **self.opt_init_args), **self.lars_init_args) return opt @@ -103,19 +117,46 @@ class ParameterProcess: parameters[i].comm_fusion = self._parameter_indices return parameters - def generate_group_params(self, parameters): + def generate_group_params(self, parameters, origin_params): """Generate group parameters.""" - decayed_params = [] - no_decayed_params = [] - for param in parameters: - if 'beta' not in param.name and 'gamma' not in param.name and 'bias' not in param.name: - decayed_params.append(param) - else: - no_decayed_params.append(param) - group_params = [{'params': decayed_params, 'weight_decay': 0.0001}, - {'params': no_decayed_params}, - {'order_params': parameters}] + origin_params_copy = origin_params + if origin_params_copy is not None: + if not isinstance(origin_params_copy, list): + origin_params_copy = list(origin_params_copy) + if not origin_params_copy: + raise ValueError("Optimizer got an empty parameter list.") + + if not isinstance(origin_params_copy[0], (dict, Parameter)): + raise TypeError("Only a list of Parameter or dict can be supported.") + + if isinstance(origin_params_copy[0], Parameter): + group_params = [{"params": parameters}] + else: + group_params = [] + params_name = [param.name for param in parameters] + new_params_count = copy.deepcopy(params_name) + for group_param in origin_params_copy: + if 'order_params' in group_param.keys(): + new_group_param = copy.deepcopy(group_param) + new_group_param['order_params'] = parameters + group_params.append(new_group_param) + continue + params_value = [] + for param in group_param['params']: + if param.name in params_name: + index = params_name.index(param.name) + params_value.append(parameters[index]) + new_params_count.remove(param.name) + new_group_param = copy.deepcopy(group_param) + new_group_param['params'] = params_value + group_params.append(new_group_param) + if new_params_count: + params_value = [] + for param in new_params_count: + index = params_name.index(param) + params_value.append(parameters[index]) + group_params.append({"params": params_value}) return group_params _gradient_accumulation_op = C.MultitypeFuncGraph("gradient_accumulation_op") diff --git a/mindspore/nn/acc/grad_freeze.py b/mindspore/nn/acc/grad_freeze.py index 635b60a141c..dd8835953ec 100644 --- a/mindspore/nn/acc/grad_freeze.py +++ b/mindspore/nn/acc/grad_freeze.py @@ -235,7 +235,8 @@ class GradientFreeze: train_para_groups = self.split_parameters_groups( network, self._param_groups) for i in range(self._param_groups): - train_para_groups[i] = self._param_processer.generate_group_params(train_para_groups[i]) + train_para_groups[i] = self._param_processer.generate_group_params(train_para_groups[i], + optimizer.init_params['params']) train_strategy = self.generate_freeze_index_sequence( self._param_groups, self._freeze_type, self._freeze_p, self._total_steps) optimizer = FreezeOpt(optimizer, train_para_groups, train_strategy) @@ -248,7 +249,7 @@ def freeze_cell(reducer_flag, network, optimizer, sens, grad, use_grad_accumulat if reducer_flag: param_processer = ParameterProcess() grad_reducers = (DistributedGradReducer(param_processer.assign_parameter_group(opt.parameters), - mean, degree, param_fusion=True) for opt in optimizer.opts) + mean, degree) for opt in optimizer.opts) freeze_nets = tuple(_TrainFreezeCell(network, sens, grad, reducer, use_grad_accumulation, opt, max_accumulation_step) for reducer, opt in zip(grad_reducers, optimizer.opts)) diff --git a/mindspore/nn/optim/ada_grad.py b/mindspore/nn/optim/ada_grad.py index d68be663549..663cbf6536f 100644 --- a/mindspore/nn/optim/ada_grad.py +++ b/mindspore/nn/optim/ada_grad.py @@ -16,6 +16,7 @@ from mindspore.ops import functional as F, composite as C, operations as P from mindspore._checkparam import Validator as validator from .optimizer import Optimizer +from .optimizer import opt_init_args_register _ada_grad_opt = C.MultitypeFuncGraph("ada_grad_opt") @@ -144,6 +145,7 @@ class Adagrad(Optimizer): >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + @opt_init_args_register def __init__(self, params, accum=0.1, learning_rate=0.001, update_slots=True, loss_scale=1.0, weight_decay=0.0): super(Adagrad, self).__init__(learning_rate, params, weight_decay, loss_scale) diff --git a/mindspore/nn/optim/adam.py b/mindspore/nn/optim/adam.py index 56a7e7cbb9e..9740731d602 100755 --- a/mindspore/nn/optim/adam.py +++ b/mindspore/nn/optim/adam.py @@ -25,6 +25,7 @@ from mindspore.common.tensor import Tensor from mindspore._checkparam import Validator as validator from mindspore._checkparam import Rel from .optimizer import Optimizer +from .optimizer import opt_init_args_register _adam_opt = C.MultitypeFuncGraph("adam_opt") _scaler_one = Tensor(1, mstype.int32) @@ -311,6 +312,7 @@ class Adam(Optimizer): >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + @opt_init_args_register def __init__(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False, use_nesterov=False, weight_decay=0.0, loss_scale=1.0): super(Adam, self).__init__(learning_rate, params, weight_decay, loss_scale) diff --git a/mindspore/nn/optim/ftrl.py b/mindspore/nn/optim/ftrl.py index aa684fb1ab8..d5c5bf328e6 100644 --- a/mindspore/nn/optim/ftrl.py +++ b/mindspore/nn/optim/ftrl.py @@ -19,6 +19,7 @@ import mindspore.common.dtype as mstype from mindspore._checkparam import Validator as validator from mindspore._checkparam import Rel from .optimizer import Optimizer, _apply_decay, _grad_scale +from .optimizer import opt_init_args_register _ftrl_opt = C.MultitypeFuncGraph("ftrl_opt") @@ -191,6 +192,8 @@ class FTRL(Optimizer): >>> loss = nn.SoftmaxCrossEntropyWithLogits() >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + + @opt_init_args_register def __init__(self, params, initial_accum=0.1, learning_rate=0.001, lr_power=-0.5, l1=0.0, l2=0.0, use_locking=False, loss_scale=1.0, weight_decay=0.0): super(FTRL, self).__init__(learning_rate, params, weight_decay, loss_scale=loss_scale) diff --git a/mindspore/nn/optim/lamb.py b/mindspore/nn/optim/lamb.py index 7a1f8a6e83e..f6ebab4f7d7 100755 --- a/mindspore/nn/optim/lamb.py +++ b/mindspore/nn/optim/lamb.py @@ -25,6 +25,7 @@ from mindspore.common.tensor import Tensor from mindspore._checkparam import Validator as validator from mindspore._checkparam import Rel from .optimizer import Optimizer +from .optimizer import opt_init_args_register from .. import layer @@ -266,6 +267,7 @@ class Lamb(Optimizer): >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + @opt_init_args_register def __init__(self, params, learning_rate, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0): super(Lamb, self).__init__(learning_rate, params, weight_decay) _check_param_value(beta1, beta2, eps, self.cls_name) diff --git a/mindspore/nn/optim/lazyadam.py b/mindspore/nn/optim/lazyadam.py index 9a26fc9d833..cd1131462cb 100644 --- a/mindspore/nn/optim/lazyadam.py +++ b/mindspore/nn/optim/lazyadam.py @@ -23,6 +23,7 @@ from mindspore.common.tensor import Tensor from mindspore._checkparam import Validator as validator from mindspore._checkparam import Rel from .optimizer import Optimizer +from .optimizer import opt_init_args_register _lazy_adam_opt = C.MultitypeFuncGraph("lazy_adam_opt") @@ -231,6 +232,7 @@ class LazyAdam(Optimizer): >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + @opt_init_args_register def __init__(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False, use_nesterov=False, weight_decay=0.0, loss_scale=1.0): super(LazyAdam, self).__init__(learning_rate, params, weight_decay, loss_scale) diff --git a/mindspore/nn/optim/proximal_ada_grad.py b/mindspore/nn/optim/proximal_ada_grad.py index f103529dd25..6a31e6246bc 100644 --- a/mindspore/nn/optim/proximal_ada_grad.py +++ b/mindspore/nn/optim/proximal_ada_grad.py @@ -18,6 +18,7 @@ from mindspore.common import Tensor import mindspore.common.dtype as mstype from mindspore._checkparam import Validator as validator from .optimizer import Optimizer +from .optimizer import opt_init_args_register _proximal_ada_grad_opt = C.MultitypeFuncGraph("proximal_ada_grad_opt") @@ -167,6 +168,7 @@ class ProximalAdagrad(Optimizer): >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + @opt_init_args_register def __init__(self, params, accum=0.1, learning_rate=0.001, l1=0.0, l2=0.0, use_locking=False, loss_scale=1.0, weight_decay=0.0): super(ProximalAdagrad, self).__init__(learning_rate, params, weight_decay, loss_scale) diff --git a/mindspore/nn/optim/rmsprop.py b/mindspore/nn/optim/rmsprop.py index 0f03a811165..72a6d08dc8f 100644 --- a/mindspore/nn/optim/rmsprop.py +++ b/mindspore/nn/optim/rmsprop.py @@ -16,6 +16,7 @@ from mindspore.ops import functional as F, composite as C, operations as P from mindspore._checkparam import Validator as validator from .optimizer import Optimizer +from .optimizer import opt_init_args_register _rmsprop_opt = C.MultitypeFuncGraph("rmsprop_opt") _centered_rmsprop_opt = C.MultitypeFuncGraph("rmsprop_opt") @@ -175,6 +176,8 @@ class RMSProp(Optimizer): >>> loss = nn.SoftmaxCrossEntropyWithLogits() >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + + @opt_init_args_register def __init__(self, params, learning_rate=0.1, decay=0.9, momentum=0.0, epsilon=1e-10, use_locking=False, centered=False, loss_scale=1.0, weight_decay=0.0): super(RMSProp, self).__init__(learning_rate, params, weight_decay, loss_scale) diff --git a/mindspore/nn/optim/sgd.py b/mindspore/nn/optim/sgd.py index d1e92f41fff..ae7b7a303e6 100755 --- a/mindspore/nn/optim/sgd.py +++ b/mindspore/nn/optim/sgd.py @@ -19,6 +19,7 @@ from mindspore.common.tensor import Tensor import mindspore.common.dtype as mstype from mindspore._checkparam import Validator as validator from .optimizer import Optimizer +from .optimizer import opt_init_args_register _sgd_opt = C.MultitypeFuncGraph("sgd_opt") @@ -133,6 +134,8 @@ class SGD(Optimizer): >>> loss = nn.SoftmaxCrossEntropyWithLogits() >>> model = Model(net, loss_fn=loss, optimizer=optim) """ + + @opt_init_args_register def __init__(self, params, learning_rate=0.1, momentum=0.0, dampening=0.0, weight_decay=0.0, nesterov=False, loss_scale=1.0): diff --git a/mindspore/nn/wrap/grad_reducer.py b/mindspore/nn/wrap/grad_reducer.py index 602f74bb44b..29836aec927 100644 --- a/mindspore/nn/wrap/grad_reducer.py +++ b/mindspore/nn/wrap/grad_reducer.py @@ -49,13 +49,25 @@ def _init_allreduce_operators(length, split_indices): def _init_allreduce_operators_by_parameters(parameters): """ initialize allreduce communication operators by parameters""" op_list = () + param_fusion = False + last_comm_fusion = None + first_parameter_flag = True for parameter in parameters: comm_fusion = parameter.comm_fusion + if first_parameter_flag: + last_comm_fusion = comm_fusion + first_parameter_flag = False + elif not param_fusion: + if comm_fusion != last_comm_fusion: + param_fusion = True + last_comm_fusion = comm_fusion op = AllReduce('sum', GlobalComm.WORLD_COMM_GROUP) op.add_prim_attr('fusion', comm_fusion) op.add_prim_attr('index', comm_fusion) op_list = op_list + (op,) - return op_list + if not param_fusion: + op_list = () + return op_list, param_fusion @reduce_opt.register("Tensor", "Bool", "Function", "Function", "Bool", "Tensor") @@ -354,7 +366,7 @@ class DistributedGradReducer(Cell): 256.0 """ - def __init__(self, parameters, mean=True, degree=None, fusion_type=1, param_fusion=False): + def __init__(self, parameters, mean=True, degree=None, fusion_type=1): super(DistributedGradReducer, self).__init__(auto_prefix=False) self.map_ = C.Map() if degree is None: @@ -371,12 +383,12 @@ class DistributedGradReducer(Cell): if is_parallel_optimizer and split_indices: self.split_fusion = True self.op_list = _init_allreduce_operators(len(parameters), split_indices) - elif param_fusion: - self.split_fusion = True - self.op_list = _init_allreduce_operators_by_parameters(parameters) else: - self.split_fusion = False - self.allreduce = AllReduce().add_prim_attr('fusion', fusion_type) + self.split_fusion = True + self.op_list, param_fusion = _init_allreduce_operators_by_parameters(parameters) + if not param_fusion: + self.split_fusion = False + self.allreduce = AllReduce().add_prim_attr('fusion', fusion_type) self.allgather = AllGather(GlobalComm.WORLD_COMM_GROUP) ps_filter = lambda x: x.is_param_ps self.ps_parameters = tuple(ps_filter(x) for x in parameters) diff --git a/mindspore/train/model.py b/mindspore/train/model.py index c1adbb0b2e0..65fa7f2b04d 100644 --- a/mindspore/train/model.py +++ b/mindspore/train/model.py @@ -139,6 +139,7 @@ class Model: self._device_number = _get_device_num() self._global_rank = _get_global_rank() self._parameter_broadcast = _get_parameter_broadcast() + self._metrics = metrics self._check_amp_level_arg(optimizer, amp_level) self._check_for_graph_cell(kwargs) @@ -175,7 +176,7 @@ class Model: def _check_kwargs(self, kwargs): for arg in kwargs: - if arg not in ['loss_scale_manager', 'keep_batchnorm_fp32', 'total_steps']: + if arg not in ['loss_scale_manager', 'keep_batchnorm_fp32']: raise ValueError(f"Unsupported arg '{arg}'") def _check_reuse_dataset(self, dataset): @@ -187,15 +188,18 @@ class Model: def _build_acc_network(self, kwargs): """Build the acc network.""" processor = acc.AutoAcc(self._acc_level, kwargs) + if processor.level not in ["O1", "O2"]: + return if self._optimizer is None: logger.warning("In acc mode, the optimizer must be defined.") return - if self._eval_network is None: - logger.warning("In acc mode, the eval_network must be defined.") + if self._eval_network is None and self._metrics is None: + logger.warning("In acc mode, the eval_network and metrics cannot be undefined at the same time.") return self._network, self._optimizer = processor.network_auto_process_train(self._network, self._optimizer) - self._eval_network = processor.network_auto_process_eval(self._eval_network) + if self._eval_network is not None: + self._eval_network = processor.network_auto_process_eval(self._eval_network) def _build_train_network(self): """Build train network""" diff --git a/model_zoo/official/cv/resnet/resnet101_imagenet2012_config.yaml b/model_zoo/official/cv/resnet/resnet101_imagenet2012_config.yaml index 7a8ff11fe68..7c6971c860c 100644 --- a/model_zoo/official/cv/resnet/resnet101_imagenet2012_config.yaml +++ b/model_zoo/official/cv/resnet/resnet101_imagenet2012_config.yaml @@ -48,6 +48,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet18_cifar10_config.yaml b/model_zoo/official/cv/resnet/resnet18_cifar10_config.yaml index 98708606c38..e164bffd506 100644 --- a/model_zoo/official/cv/resnet/resnet18_cifar10_config.yaml +++ b/model_zoo/official/cv/resnet/resnet18_cifar10_config.yaml @@ -48,6 +48,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet18_imagenet2012_config.yaml b/model_zoo/official/cv/resnet/resnet18_imagenet2012_config.yaml index 97a800959f1..92c66f238a2 100644 --- a/model_zoo/official/cv/resnet/resnet18_imagenet2012_config.yaml +++ b/model_zoo/official/cv/resnet/resnet18_imagenet2012_config.yaml @@ -50,6 +50,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet34_imagenet2012_config.yaml b/model_zoo/official/cv/resnet/resnet34_imagenet2012_config.yaml index fe8dd627f58..5b4b0493dfa 100644 --- a/model_zoo/official/cv/resnet/resnet34_imagenet2012_config.yaml +++ b/model_zoo/official/cv/resnet/resnet34_imagenet2012_config.yaml @@ -50,6 +50,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet50_cifar10_config.yaml b/model_zoo/official/cv/resnet/resnet50_cifar10_config.yaml index e568d5f0ef7..428bfcd2b40 100644 --- a/model_zoo/official/cv/resnet/resnet50_cifar10_config.yaml +++ b/model_zoo/official/cv/resnet/resnet50_cifar10_config.yaml @@ -48,6 +48,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet50_imagenet2012_Acc_config.yaml b/model_zoo/official/cv/resnet/resnet50_imagenet2012_Acc_config.yaml new file mode 100644 index 00000000000..6b0a8832b65 --- /dev/null +++ b/model_zoo/official/cv/resnet/resnet50_imagenet2012_Acc_config.yaml @@ -0,0 +1,78 @@ +# Builtin Configurations(DO NOT CHANGE THESE CONFIGURATIONS unless you know exactly what you are doing) +enable_modelarts: False +# Url for modelarts +data_url: "" +train_url: "" +checkpoint_url: "" +# Path for local +run_distribute: False +enable_profiling: False +data_path: "/cache/data" +output_path: "/cache/train" +load_path: "/cache/checkpoint_path/" +device_target: "Ascend" +checkpoint_path: "./checkpoint/" +checkpoint_file_path: "" + +# ============================================================================== +# Training options +optimizer: "Momentum" +infer_label: "" +class_num: 1001 +batch_size: 256 +loss_scale: 1024 +momentum: 0.9 +weight_decay: 0.0001 +epoch_size: 90 +pretrain_epoch_size: 0 +save_checkpoint: True +save_checkpoint_epochs: 5 +keep_checkpoint_max: 10 +warmup_epochs: 0 +lr_decay_mode: "linear" +use_label_smooth: True +label_smooth_factor: 0.1 +lr_init: 0 +lr_max: 0.8 +lr_end: 0.0 + +net_name: "resnet50" +dataset: "imagenet2012" +device_num: 1 +pre_trained: "" +run_eval: False +eval_dataset_path: "" +parameter_server: False +filter_weight: False +save_best_ckpt: True +eval_start_epoch: 40 +eval_interval: 1 +enable_cache: False +cache_session_id: "" +mode_name: "GRAPH" +acc_mode: "O1" + +# Export options +device_id: 0 +width: 224 +height: 224 +file_name: "resnet50" +file_format: "AIR" +ckpt_file: "" +network_dataset: "resnet50_imagenet2012" + +--- +# Help description for each configuration +enable_modelarts: "Whether training on modelarts, default: False" +data_url: "Dataset url for obs" +checkpoint_url: "The location of checkpoint for obs" +data_path: "Dataset path for local" +output_path: "Training output path for local" +load_path: "The location of checkpoint for obs" +device_target: "Target device type, available: [Ascend, GPU, CPU]" +enable_profiling: "Whether enable profiling while training, default: False" +num_classes: "Class for dataset" +batch_size: "Batch size for training and evaluation" +epoch_size: "Total training epochs." +checkpoint_path: "The location of the checkpoint file." +checkpoint_file_path: "The location of the checkpoint file." diff --git a/model_zoo/official/cv/resnet/resnet50_imagenet2012_GPU_config.yaml b/model_zoo/official/cv/resnet/resnet50_imagenet2012_GPU_config.yaml index 603e4a9ae32..0bbce0d7d80 100644 --- a/model_zoo/official/cv/resnet/resnet50_imagenet2012_GPU_config.yaml +++ b/model_zoo/official/cv/resnet/resnet50_imagenet2012_GPU_config.yaml @@ -51,6 +51,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet50_imagenet2012_Thor_config.yaml b/model_zoo/official/cv/resnet/resnet50_imagenet2012_Thor_config.yaml index 73528f2de33..29b8997e2b7 100644 --- a/model_zoo/official/cv/resnet/resnet50_imagenet2012_Thor_config.yaml +++ b/model_zoo/official/cv/resnet/resnet50_imagenet2012_Thor_config.yaml @@ -51,6 +51,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet50_imagenet2012_config.yaml b/model_zoo/official/cv/resnet/resnet50_imagenet2012_config.yaml index 96c5566d448..b86f26f05fd 100644 --- a/model_zoo/official/cv/resnet/resnet50_imagenet2012_config.yaml +++ b/model_zoo/official/cv/resnet/resnet50_imagenet2012_config.yaml @@ -50,6 +50,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/resnet_benchmark_GPU.yaml b/model_zoo/official/cv/resnet/resnet_benchmark_GPU.yaml index c7c0455b4d6..8c341a0bbd3 100644 --- a/model_zoo/official/cv/resnet/resnet_benchmark_GPU.yaml +++ b/model_zoo/official/cv/resnet/resnet_benchmark_GPU.yaml @@ -25,6 +25,7 @@ eval: False save_ckpt: False mode_name: "GRAPH" dtype: "fp16" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/se-resnet50_imagenet2012_config.yaml b/model_zoo/official/cv/resnet/se-resnet50_imagenet2012_config.yaml index a6e01b850cf..f313c1098e4 100644 --- a/model_zoo/official/cv/resnet/se-resnet50_imagenet2012_config.yaml +++ b/model_zoo/official/cv/resnet/se-resnet50_imagenet2012_config.yaml @@ -51,6 +51,7 @@ eval_interval: 1 enable_cache: False cache_session_id: "" mode_name: "GRAPH" +acc_mode: "O0" # Export options device_id: 0 diff --git a/model_zoo/official/cv/resnet/src/dataset.py b/model_zoo/official/cv/resnet/src/dataset.py index 13d76701f7f..34ab2869a6b 100755 --- a/model_zoo/official/cv/resnet/src/dataset.py +++ b/model_zoo/official/cv/resnet/src/dataset.py @@ -124,9 +124,9 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target= device_num = 1 if device_num == 1: - data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=8, shuffle=True) + data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=12, shuffle=True) else: - data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=8, shuffle=True, + data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=12, shuffle=True, num_shards=device_num, shard_id=rank_id) image_size = 224 @@ -152,7 +152,7 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target= type_cast_op = C2.TypeCast(mstype.int32) - data_set = data_set.map(operations=trans, input_columns="image", num_parallel_workers=8) + data_set = data_set.map(operations=trans, input_columns="image", num_parallel_workers=12) # only enable cache for eval if do_train: enable_cache = False @@ -160,10 +160,10 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target= if not cache_session_id: raise ValueError("A cache session_id must be provided to use cache.") eval_cache = ds.DatasetCache(session_id=int(cache_session_id), size=0) - data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=8, + data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=12, cache=eval_cache) else: - data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=8) + data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=12) # apply batch operations data_set = data_set.batch(batch_size, drop_remainder=True) diff --git a/model_zoo/official/cv/resnet/train.py b/model_zoo/official/cv/resnet/train.py index e7d675b4c51..ecdc24325db 100755 --- a/model_zoo/official/cv/resnet/train.py +++ b/model_zoo/official/cv/resnet/train.py @@ -105,7 +105,8 @@ def set_parameter(): gradients_mean=True) set_algo_parameters(elementwise_op_strategy_follow=True) if config.net_name == "resnet50" or config.net_name == "se-resnet50": - context.set_auto_parallel_context(all_reduce_fusion_config=[85, 160]) + if config.acc_mode not in ["O1", "O2"]: + context.set_auto_parallel_context(all_reduce_fusion_config=[85, 160]) elif config.net_name == "resnet101": context.set_auto_parallel_context(all_reduce_fusion_config=[80, 210, 313]) init() @@ -228,7 +229,8 @@ def train_net(): model = Model(net, loss_fn=loss, optimizer=opt, metrics=metrics, eval_network=dist_eval_network) else: model = Model(net, loss_fn=loss, optimizer=opt, loss_scale_manager=loss_scale, metrics=metrics, - amp_level="O2", keep_batchnorm_fp32=False, eval_network=dist_eval_network) + amp_level="O2", acc_level=config.acc_mode, keep_batchnorm_fp32=False, + eval_network=dist_eval_network) if config.optimizer == "Thor" and config.dataset == "imagenet2012": from src.lr_generator import get_thor_damping