diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/adam_apply_one_with_decay_rule.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/adam_apply_one_with_decay_rule.cc index b2e4f40e03d..5f6642f8e25 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/adam_apply_one_with_decay_rule.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/adam_apply_one_with_decay_rule.cc @@ -279,6 +279,7 @@ const AnfNodePtr AdamApplyOneWithDecayRule::Process(const FuncGraphPtr &graph, c return nullptr; } auto sub0 = node; + constexpr size_t kOutputIndex2 = 2; if (AnfAlgo::CheckPrimitiveType(node, prim::kPrimDepend)) { auto iter_sub0 = (*equiv).find(sub0_var_); if (iter_sub0 == (*equiv).end()) { @@ -327,7 +328,7 @@ const AnfNodePtr AdamApplyOneWithDecayRule::Process(const FuncGraphPtr &graph, c MS_EXCEPTION_IF_NULL(manager); (void)manager->Replace(add1, fusion_node_outputs[0]); (void)manager->Replace(add0, fusion_node_outputs[1]); - return fusion_node_outputs[2]; + return fusion_node_outputs[kOutputIndex2]; } } // namespace opt } // namespace mindspore diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/batchnormgrad_to_bninfergrad.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/batchnormgrad_to_bninfergrad.cc index cddc3b6a568..f12a1dfe8df 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/batchnormgrad_to_bninfergrad.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/batchnormgrad_to_bninfergrad.cc @@ -31,9 +31,12 @@ CNodePtr CreateBNInferGrad(const FuncGraphPtr &graph, const CNodePtr &batchnormg MS_EXCEPTION_IF_NULL(batchnormgrad); auto prim = std::make_shared(kBNInferGradOpName); std::vector inputs = {NewValueNode(prim)}; - inputs.push_back(batchnormgrad->input(1)); - inputs.push_back(batchnormgrad->input(3)); - inputs.push_back(batchnormgrad->input(5)); + constexpr size_t kDBatchMean = 1; + constexpr size_t kInputX = 3; + constexpr size_t kBatchStd = 5; + inputs.push_back(batchnormgrad->input(kDBatchMean)); + inputs.push_back(batchnormgrad->input(kInputX)); + inputs.push_back(batchnormgrad->input(kBatchStd)); auto new_node = graph->NewCNode(inputs); MS_EXCEPTION_IF_NULL(new_node); new_node->set_scope(batchnormgrad->scope()); diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/derelu_fusion.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/derelu_fusion.cc index cd032a207e0..c8ccc2198b6 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/derelu_fusion.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/derelu_fusion.cc @@ -30,8 +30,9 @@ const size_t kReluV2OutputNum = 2; CNodePtr GetRelu(const CNodePtr &relu_grad) { MS_EXCEPTION_IF_NULL(relu_grad); + constexpr size_t kReluIndex2 = 2; CheckCNodeInputSize(relu_grad, kReluGradInputTensorNum); - auto relu_anf = relu_grad->input(2); + auto relu_anf = relu_grad->input(kReluIndex2); MS_EXCEPTION_IF_NULL(relu_anf); return relu_anf->cast(); } @@ -40,7 +41,7 @@ CNodePtr CreateReluV2(const FuncGraphPtr &graph, const CNodePtr &relu) { MS_EXCEPTION_IF_NULL(graph); MS_EXCEPTION_IF_NULL(relu); CheckCNodeInputSize(relu, kReluInputTensorNum); - + constexpr auto kMaskShapeSize = 4; auto prim = std::make_shared(kReluV2OpName); std::vector inputs = {NewValueNode(prim), relu->input(1)}; auto new_node = graph->NewCNode(inputs); @@ -53,17 +54,24 @@ CNodePtr CreateReluV2(const FuncGraphPtr &graph, const CNodePtr &relu) { return nullptr; } std::vector mask_shape = AnfAlgo::GetOutputInferShape(relu, 0); - if (mask_shape.size() != 4) { + if (mask_shape.size() != kMaskShapeSize) { MS_LOG(DEBUG) << "relu's infer shape size not equal 4"; return nullptr; } auto input_dtype = AnfAlgo::GetPrevNodeOutputInferDataType(relu, 0); + constexpr auto kMultiplierInt8 = 31; + constexpr auto kDivisorInt8 = 32; + constexpr auto kMultiplier = 15; + constexpr auto kDivisor = 16; + constexpr auto kThirdShapeInt8 = 4; + constexpr auto kThirdShape = 2; + if (input_dtype == kNumberTypeUInt8 || input_dtype == kNumberTypeInt8) { - mask_shape[1] = (mask_shape[1] + 31) / 32; - mask_shape.push_back(4); + mask_shape[1] = (mask_shape[1] + kMultiplierInt8) / kDivisorInt8; + mask_shape.push_back(kThirdShapeInt8); } else { - mask_shape[1] = (mask_shape[1] + 15) / 16; - mask_shape.push_back(2); + mask_shape[1] = (mask_shape[1] + kMultiplier) / kDivisor; + mask_shape.push_back(kThirdShape); } auto types = {AnfAlgo::GetOutputInferDataType(relu, 0), mask_dtype}; diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/fused_batch_norm_fusion.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/fused_batch_norm_fusion.cc index 82708121d60..32479f5d871 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/fused_batch_norm_fusion.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/fused_batch_norm_fusion.cc @@ -138,8 +138,11 @@ void FusedBatchNormFusion::GetBNTrainingUpdateAbstractList(const EquivPtr &equiv auto variable_input1 = GetAnfNodeByVar(equiv, variable_input1_var_); MS_EXCEPTION_IF_NULL(variable_input0); MS_EXCEPTION_IF_NULL(variable_input1); - *abstract_list = {bn_abstract_tuple->elements()[0], variable_input0->abstract(), variable_input1->abstract(), - bn_abstract_tuple->elements()[1], bn_abstract_tuple->elements()[2]}; + constexpr size_t kElements0 = 0; + constexpr size_t kElements1 = 1; + constexpr size_t kElements2 = 2; + *abstract_list = {bn_abstract_tuple->elements()[kElements0], variable_input0->abstract(), variable_input1->abstract(), + bn_abstract_tuple->elements()[kElements1], bn_abstract_tuple->elements()[kElements2]}; } AnfNodePtr FusedBatchNormFusion::CreateBNTrainingUpdate( diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_rule.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_rule.cc index c487b6a3cb1..280b5bc1961 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_rule.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_rule.cc @@ -30,11 +30,12 @@ bool LambNextMVRule::IsRuleMatched(const FuncGraphPtr &func_graph, const AnfNode MS_EXCEPTION_IF_NULL(equiv); auto real_div0 = GetAnfNodeByVar(equiv, real_div0_var_); auto real_div2 = GetAnfNodeByVar(equiv, real_div2_var_); + constexpr size_t kRealDiv0Size = 2; auto manager = func_graph->manager(); MS_EXCEPTION_IF_NULL(manager); auto &users = manager->node_users(); - if (users.find(real_div0) == users.end() || users[real_div0].size() < 2) { + if (users.find(real_div0) == users.end() || users[real_div0].size() < kRealDiv0Size) { return false; } AnfNodeIndexSet real_div0_outputs = users[real_div0]; @@ -60,6 +61,10 @@ AnfNodePtr LambNextMVRule::CreateLambNextMVNode(const FuncGraphPtr &func_graph, const EquivPtr &equiv) const { MS_EXCEPTION_IF_NULL(func_graph); auto prim = std::make_shared(kLambNextMVOpName); + constexpr size_t kOutputsIndex1 = 1; + constexpr size_t kOutputsIndex2 = 2; + constexpr size_t kOutputsIndex3 = 3; + std::vector lamb_next_mv_rule_inputs = {NewValueNode(prim)}; lamb_next_mv_rule_inputs.push_back(utils::cast((*equiv)[input0_])); lamb_next_mv_rule_inputs.push_back(utils::cast((*equiv)[input1_])); @@ -91,9 +96,9 @@ AnfNodePtr LambNextMVRule::CreateLambNextMVNode(const FuncGraphPtr &func_graph, auto manager = func_graph->manager(); MS_EXCEPTION_IF_NULL(manager); - (void)manager->Replace(old_pattern_outputs[1], lamb_next_mv_rule_outputs[1]); - (void)manager->Replace(old_pattern_outputs[2], lamb_next_mv_rule_outputs[2]); - (void)manager->Replace(old_pattern_outputs[3], lamb_next_mv_rule_outputs[3]); + (void)manager->Replace(old_pattern_outputs[kOutputsIndex1], lamb_next_mv_rule_outputs[kOutputsIndex1]); + (void)manager->Replace(old_pattern_outputs[kOutputsIndex2], lamb_next_mv_rule_outputs[kOutputsIndex2]); + (void)manager->Replace(old_pattern_outputs[kOutputsIndex3], lamb_next_mv_rule_outputs[kOutputsIndex3]); return lamb_next_mv_rule_outputs[0]; } diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_rule.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_rule.cc index f1a840adada..9599321aefd 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_rule.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_rule.cc @@ -19,6 +19,16 @@ #include "frontend/optimizer/opt.h" #include "utils/trace_base.h" namespace mindspore { +namespace { +constexpr size_t kZeroIndex = 0; +constexpr size_t kFirstIndex = 1; +constexpr size_t kSecondIndex = 2; +constexpr size_t kThirdIndex = 3; +constexpr size_t kFourthIndex = 4; +constexpr size_t kFifthIndex = 5; +constexpr size_t kSixthIndex = 6; +} // namespace + namespace opt { AnfNodePtr LambNextMVWithDecayRule::GetLambNextMVWithDecayOutput(const FuncGraphPtr &func_graph, const AnfNodePtr &new_node, const AnfNodePtr &add3, @@ -136,18 +146,18 @@ const BaseRef LambNextMVWithDecayRuleCond1::DefinePattern() const { MS_EXCEPTION_IF_NULL(prim_sqrt); const auto prim_deal_div = std::make_shared(kRealDivOpName); MS_EXCEPTION_IF_NULL(prim_deal_div); - VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[1], constant_mul_input_vars_[2]}); - VectorRef mul3 = VectorRef({prim::kPrimMul, input_vars_[0], constant_mul_input_vars_[3]}); + VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[kFirstIndex], constant_mul_input_vars_[kSecondIndex]}); + VectorRef mul3 = VectorRef({prim::kPrimMul, input_vars_[kZeroIndex], constant_mul_input_vars_[kThirdIndex]}); VectorRef add1 = VectorRef({add1_var_, mul2, mul3}); - VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]}); + VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]}); VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1}); VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_}); - VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[4], constant_mul_input_vars_[0]}); - VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[3], constant_mul_input_vars_[1]}); + VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[kFourthIndex], constant_mul_input_vars_[kZeroIndex]}); + VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[kThirdIndex], constant_mul_input_vars_[kFirstIndex]}); VectorRef add0 = VectorRef({add0_var_, mul0, mul1}); - VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]}); + VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]}); VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4}); - VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]}); + VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]}); VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4}); return add5; } @@ -177,18 +187,18 @@ const BaseRef LambNextMVWithDecayRuleCond2::DefinePattern() const { MS_EXCEPTION_IF_NULL(prim_sqrt); const auto prim_deal_div = std::make_shared(kRealDivOpName); MS_EXCEPTION_IF_NULL(prim_deal_div); - VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[2], input_vars_[1]}); - VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]}); + VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kSecondIndex], input_vars_[kFirstIndex]}); + VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]}); VectorRef add1 = VectorRef({add1_var_, mul2, mul3}); - VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]}); + VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]}); VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1}); VectorRef add4 = VectorRef({prim::kPrimAdd, constant_add2_y_, sqrt1}); - VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[0], input_vars_[4]}); - VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[1], input_vars_[3]}); + VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kZeroIndex], input_vars_[kFourthIndex]}); + VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kFirstIndex], input_vars_[kThirdIndex]}); VectorRef add0 = VectorRef({add0_var_, mul0, mul1}); - VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]}); + VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]}); VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4}); - VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]}); + VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]}); VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4}); return add5; } @@ -218,18 +228,18 @@ const BaseRef LambNextMVWithDecayRuleCond3::DefinePattern() const { MS_EXCEPTION_IF_NULL(prim_sqrt); const auto prim_deal_div = std::make_shared(kRealDivOpName); MS_EXCEPTION_IF_NULL(prim_deal_div); - VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[1], constant_mul_input_vars_[2]}); - VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]}); + VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[kFirstIndex], constant_mul_input_vars_[kSecondIndex]}); + VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]}); VectorRef add1 = VectorRef({add1_var_, mul2, mul3}); - VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]}); + VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]}); VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1}); VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_}); - VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[4], constant_mul_input_vars_[0]}); - VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[3], constant_mul_input_vars_[1]}); + VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[kFourthIndex], constant_mul_input_vars_[kZeroIndex]}); + VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[kThirdIndex], constant_mul_input_vars_[kFirstIndex]}); VectorRef add0 = VectorRef({add0_var_, mul0, mul1}); - VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]}); + VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]}); VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4}); - VectorRef mul4 = VectorRef({mul4_var_, input_vars_[6], constant_mul_input_vars_[4]}); + VectorRef mul4 = VectorRef({mul4_var_, input_vars_[kSixthIndex], constant_mul_input_vars_[kFourthIndex]}); VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4}); return add5; } @@ -260,18 +270,18 @@ const BaseRef LambNextMVWithDecayRuleCond4::DefinePattern() const { MS_EXCEPTION_IF_NULL(prim_sqrt); const auto prim_deal_div = std::make_shared(kRealDivOpName); MS_EXCEPTION_IF_NULL(prim_deal_div); - VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[2], input_vars_[1]}); - VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]}); + VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kSecondIndex], input_vars_[kFirstIndex]}); + VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]}); VectorRef add1 = VectorRef({add1_var_, mul2, mul3}); - VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]}); + VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]}); VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1}); VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_}); - VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[0], input_vars_[4]}); - VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[1], input_vars_[3]}); + VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kZeroIndex], input_vars_[kFourthIndex]}); + VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kFirstIndex], input_vars_[kThirdIndex]}); VectorRef add0 = VectorRef({add0_var_, mul0, mul1}); - VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]}); + VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]}); VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4}); - VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]}); + VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]}); VectorRef add5 = VectorRef({prim::kPrimAdd, real_div4, mul4}); return add5; } diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_v1_rule.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_v1_rule.cc index 5c31af153b2..1cb47a39416 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_v1_rule.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_next_mv_with_decay_v1_rule.cc @@ -26,26 +26,31 @@ namespace mindspore { namespace opt { namespace { +constexpr auto kFirstIndex1 = 1; +constexpr auto kSecondIndex2 = 2; +constexpr auto kThirdIndex3 = 3; + std::tuple GetSharedNodes(const AnfNodePtr &node) { MS_EXCEPTION_IF_NULL(node); auto add3 = node->cast(); MS_EXCEPTION_IF_NULL(add3); CheckCNodeInputSize(add3, kAddInputTensorNum); - auto real_div2_anf = add3->input(1); + auto real_div2_anf = add3->input(kFirstIndex1); MS_EXCEPTION_IF_NULL(real_div2_anf); auto real_div2 = real_div2_anf->cast(); MS_EXCEPTION_IF_NULL(real_div2); CheckCNodeInputSize(real_div2, kRealDivInputTensorNum); - auto sqrt0_anf = real_div2->input(2); + auto sqrt0_anf = real_div2->input(kSecondIndex2); MS_EXCEPTION_IF_NULL(sqrt0_anf); auto sqrt0 = sqrt0_anf->cast(); MS_EXCEPTION_IF_NULL(sqrt0); CheckCNodeInputSize(sqrt0, kSqrtInputTensorNum); - auto add2_anf = sqrt0->input(1); + auto add2_anf = sqrt0->input(kFirstIndex1); MS_EXCEPTION_IF_NULL(add2_anf); auto add2 = add2_anf->cast(); CheckCNodeInputSize(add2, kAddInputTensorNum); - return std::make_tuple(add3->input(2), real_div2->input(1), add2->input(1), add2->input(2)); + return std::make_tuple(add3->input(kSecondIndex2), real_div2->input(kFirstIndex1), add2->input(kFirstIndex1), + add2->input(kSecondIndex2)); } bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfNodePtr &real_div0, @@ -57,7 +62,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN if (AnfAlgo::GetCNodeName(add5) != prim::kPrimAdd->name() || AnfAlgo::GetInputTensorNum(add5) != kAddInputTensorNum) { return false; } - auto real_div4_anf = add5->input(1); + auto real_div4_anf = add5->input(kFirstIndex1); if (real_div4_anf == nullptr || !real_div4_anf->isa()) { return false; } @@ -66,7 +71,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN AnfAlgo::GetInputTensorNum(real_div4) != kRealDivInputTensorNum) { return false; } - auto add4_anf = real_div4->input(2); + auto add4_anf = real_div4->input(kSecondIndex2); if (add4_anf == nullptr || !add4_anf->isa()) { return false; } @@ -74,7 +79,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN if (AnfAlgo::GetCNodeName(add4) != prim::kPrimAdd->name() || AnfAlgo::GetInputTensorNum(add4) != kAddInputTensorNum) { return false; } - auto sqrt1_anf = add4->input(1); + auto sqrt1_anf = add4->input(kFirstIndex1); if (sqrt1_anf == nullptr || !sqrt1_anf->isa()) { return false; } @@ -82,8 +87,8 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN if (AnfAlgo::GetCNodeName(sqrt1) != kSqrtOpName || AnfAlgo::GetInputTensorNum(sqrt1) != kSqrtInputTensorNum) { return false; } - return add5->input(2) == mul4 && real_div4->input(1) == real_div0 && sqrt1->input(1) == real_div1 && - *add4->input(2) == *add2_y; + return add5->input(kSecondIndex2) == mul4 && real_div4->input(kFirstIndex1) == real_div0 && + sqrt1->input(kFirstIndex1) == real_div1 && *add4->input(kSecondIndex2) == *add2_y; } std::tuple GetAdd0Add1Nodes(const AnfNodePtr &real_div0_anf, const AnfNodePtr &real_div1_anf) { @@ -188,10 +193,9 @@ const AnfNodePtr LambNextMVWithDecayV1Rule::Process(const FuncGraphPtr &func_gra MS_LOG(EXCEPTION) << "create multiple outputs for fusion node fail!" << " trace: " << trace::DumpSourceLines(node); } - - (void)manager->Replace(add0, fusion_node_outputs[1]); - (void)manager->Replace(add1, fusion_node_outputs[2]); - (void)manager->Replace(add5, fusion_node_outputs[3]); + (void)manager->Replace(add0, fusion_node_outputs[kFirstIndex1]); + (void)manager->Replace(add1, fusion_node_outputs[kSecondIndex2]); + (void)manager->Replace(add5, fusion_node_outputs[kThirdIndex3]); return fusion_node_outputs[0]; } } // namespace opt diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_update_with_lr_v2.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_update_with_lr_v2.cc index 11a8e0b0025..c37e561cd1d 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_update_with_lr_v2.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/lamb_update_with_lr_v2.cc @@ -25,16 +25,23 @@ namespace opt { const BaseRef LambUpdateWithLrV2::DefinePattern() const { const auto prim_greater = std::make_shared(kGreaterOpName); const auto prim_deal_div = std::make_shared(kRealDivOpName); + const size_t kZeroIndex = 0; + const size_t kFirstIndex = 1; + const size_t kSecondIndex = 2; + const size_t kThirdIndex = 3; + const size_t kFourthIndex = 4; + const size_t kFifthIndex = 5; + const size_t kSixthIndex = 6; - VectorRef greater0({prim_greater, input_varptr_[0], input_varptr_[5]}); - VectorRef greater1({prim_greater, input_varptr_[1], input_varptr_[5]}); - VectorRef real_div0({prim_deal_div, input_varptr_[0], input_varptr_[1]}); - VectorRef select0({prim::kPrimSelect, greater1, real_div0, input_varptr_[6]}); - VectorRef select1({prim::kPrimSelect, greater0, select0, input_varptr_[6]}); - VectorRef mul0({prim::kPrimMul, select1, input_varptr_[2]}); - VectorRef mul1({prim::kPrimMul, mul0, input_varptr_[3]}); + VectorRef greater0({prim_greater, input_varptr_[kZeroIndex], input_varptr_[kFifthIndex]}); + VectorRef greater1({prim_greater, input_varptr_[kFirstIndex], input_varptr_[kFifthIndex]}); + VectorRef real_div0({prim_deal_div, input_varptr_[kZeroIndex], input_varptr_[kFirstIndex]}); + VectorRef select0({prim::kPrimSelect, greater1, real_div0, input_varptr_[kSixthIndex]}); + VectorRef select1({prim::kPrimSelect, greater0, select0, input_varptr_[kSixthIndex]}); + VectorRef mul0({prim::kPrimMul, select1, input_varptr_[kSecondIndex]}); + VectorRef mul1({prim::kPrimMul, mul0, input_varptr_[kThirdIndex]}); - return VectorRef({prim::kPrimSub, input_varptr_[4], mul1}); + return VectorRef({prim::kPrimSub, input_varptr_[kFourthIndex], mul1}); } const AnfNodePtr LambUpdateWithLrV2::Process(const FuncGraphPtr &func_graph, const AnfNodePtr &node, diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/momentum_lossscale_fusion.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/momentum_lossscale_fusion.cc index 4905b946d32..a08e91377a5 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/momentum_lossscale_fusion.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/momentum_lossscale_fusion.cc @@ -53,9 +53,15 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(node); auto cnode = node->cast(); + constexpr size_t kFirstIndex = 1; + constexpr size_t kSecondIndex = 2; + constexpr size_t kThirdIndex = 3; + constexpr size_t kFourthIndex = 4; + constexpr size_t kFifthIndex = 5; + constexpr size_t kSixthIndex = 6; MS_EXCEPTION_IF_NULL(cnode); CheckCNodeInputSize(cnode, kApplyMomentumInputTensorNum); - AnfNodePtr mul = cnode->input(4); + AnfNodePtr mul = cnode->input(kFourthIndex); MS_EXCEPTION_IF_NULL(mul); auto mul_cnode = mul->cast(); MS_EXCEPTION_IF_NULL(mul_cnode); @@ -74,13 +80,14 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph } auto new_prim = std::make_shared(kFusedMulApplyMomentumOpName); auto depend_prim = NewValueNode(prim::kPrimDepend); - auto depend = func_graph->NewCNode({depend_prim, cnode->input(5), cnode->input(6)}); // depend on monad - depend->set_abstract(cnode->input(5)->abstract()); - depend->set_scope(cnode->input(5)->scope()); + auto depend = + func_graph->NewCNode({depend_prim, cnode->input(kFifthIndex), cnode->input(kSixthIndex)}); // depend on monad + depend->set_abstract(cnode->input(kFifthIndex)->abstract()); + depend->set_scope(cnode->input(kFifthIndex)->scope()); std::vector new_node_inputs{NewValueNode(new_prim), - cnode->input(1), - cnode->input(2), - cnode->input(3), + cnode->input(kFirstIndex), + cnode->input(kSecondIndex), + cnode->input(kThirdIndex), mul_cnode->input(kMulInputTensorNum + 1 - value_node_index), depend, mul_cnode->input(value_node_index)}; @@ -88,7 +95,7 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph MS_EXCEPTION_IF_NULL(new_node); AnfAlgo::CopyNodeAttrs(node, new_node); auto input_names_value = AnfAlgo::GetNodeAttr>(new_node, kAttrInputNames); - input_names_value[3] = "x1"; + input_names_value[kThirdIndex] = "x1"; input_names_value.emplace_back("x2"); AnfAlgo::SetNodeAttr(kAttrInputNames, MakeValue(input_names_value), new_node); new_node->set_abstract(node->abstract()); diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/mul_add_fusion.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/mul_add_fusion.cc index 85599b69756..e456254c3cf 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/mul_add_fusion.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/mul_add_fusion.cc @@ -25,7 +25,7 @@ namespace mindspore { namespace opt { namespace { -bool GetMul(const FuncGraphPtr &graph, const CNodePtr &add, CNodePtr *mul, size_t *mul_index) { +bool GetMul(const FuncGraphPtr &graph, const CNodePtr &add, CNodePtr *const mul, size_t *const mul_index) { MS_EXCEPTION_IF_NULL(graph); MS_EXCEPTION_IF_NULL(add); diff --git a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/reshape_transpose_fusion.cc b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/reshape_transpose_fusion.cc index 2e731dcf187..c9ebdab2382 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/reshape_transpose_fusion.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/ir_fusion/reshape_transpose_fusion.cc @@ -28,10 +28,13 @@ bool CheckShapeDimInfo(const std::vector &shape) { if (shape.empty()) { return false; } - if (shape.size() == 1 && shape[0] % kCubeSize != 0) { + constexpr auto kShapeSize1 = 1; + constexpr auto kShapeSize2 = 2; + if (shape.size() == kShapeSize1 && shape[0] % kCubeSize != 0) { return false; } - return !(shape.size() >= 2 && (shape[shape.size() - 1] % kCubeSize != 0 || shape[shape.size() - 2] % kCubeSize != 0)); + return !(shape.size() >= kShapeSize2 && + (shape[shape.size() - 1] % kCubeSize != 0 || shape[shape.size() - kShapeSize2] % kCubeSize != 0)); } } // namespace