diff --git a/mindspore/ccsrc/pre_activate/ascend/ascend_backend_optimization.cc b/mindspore/ccsrc/pre_activate/ascend/ascend_backend_optimization.cc index 4011e0fe140..effec812e4b 100644 --- a/mindspore/ccsrc/pre_activate/ascend/ascend_backend_optimization.cc +++ b/mindspore/ccsrc/pre_activate/ascend/ascend_backend_optimization.cc @@ -81,6 +81,7 @@ #include "pre_activate/ascend/enhancer/getnext_memcpy_elimination.h" #include "pre_activate/ascend/ir_fission/addn_fission.h" #include "pre_activate/ascend/enhancer/insert_memcpy_async_for_getnext.h" +#include "pre_activate/ascend/ir_fission/batch_norm_grad_infer_fission.h" #include "utils/context/ms_context.h" #include "utils/config_manager.h" #include "debug/anf_ir_dump.h" @@ -116,6 +117,7 @@ void AddAscendBackendOptionalIRFusion(PassManager *ir_fusion_pm) { ir_fusion_pm->AddPass(std::make_shared()); ir_fusion_pm->AddPass(std::make_shared()); ir_fusion_pm->AddPass(std::make_shared()); + ir_fusion_pm->AddPass(std::make_shared()); } } // namespace diff --git a/mindspore/ccsrc/pre_activate/ascend/ir_fission/batch_norm_grad_infer_fission.cc b/mindspore/ccsrc/pre_activate/ascend/ir_fission/batch_norm_grad_infer_fission.cc index e1399281343..5e411116607 100644 --- a/mindspore/ccsrc/pre_activate/ascend/ir_fission/batch_norm_grad_infer_fission.cc +++ b/mindspore/ccsrc/pre_activate/ascend/ir_fission/batch_norm_grad_infer_fission.cc @@ -34,6 +34,9 @@ bool CheckOutputsIndex(const FuncGraphPtr &func_graph, const AnfNodePtr &node) { for (const auto &node_index : manager->node_users()[node]) { AnfNodePtr output = node_index.first; MS_EXCEPTION_IF_NULL(output); + if (!IsPrimitiveCNode(output, prim::kPrimTupleGetItem)) { + continue; + } auto tuple_getiterm_cnode = output->cast(); MS_EXCEPTION_IF_NULL(tuple_getiterm_cnode); auto index_node = tuple_getiterm_cnode->input(kInputNodeOutputIndexInTupleGetItem); diff --git a/mindspore/ccsrc/pre_activate/ascend/ir_fusion/fused_batch_norm_fusion.cc b/mindspore/ccsrc/pre_activate/ascend/ir_fusion/fused_batch_norm_fusion.cc index 8e48780936b..03428e63578 100644 --- a/mindspore/ccsrc/pre_activate/ascend/ir_fusion/fused_batch_norm_fusion.cc +++ b/mindspore/ccsrc/pre_activate/ascend/ir_fusion/fused_batch_norm_fusion.cc @@ -274,6 +274,9 @@ const AnfNodePtr FusedBatchNormFusion::Process(const FuncGraphPtr &func_graph, c MS_EXCEPTION_IF_NULL(manager); for (const auto &output : bn_outputs) { MS_EXCEPTION_IF_NULL(output); + if (!IsPrimitiveCNode(output, prim::kPrimTupleGetItem)) { + continue; + } auto tuple_getitem_cnode = output->cast(); MS_EXCEPTION_IF_NULL(tuple_getitem_cnode); AnfNodePtr index_node = tuple_getitem_cnode->input(kInputNodeOutputIndexInTupleGetItem); diff --git a/mindspore/ccsrc/pre_activate/ascend/ir_fusion/momentum_lossscale_fusion.cc b/mindspore/ccsrc/pre_activate/ascend/ir_fusion/momentum_lossscale_fusion.cc index 8833e75c761..6b751873d68 100644 --- a/mindspore/ccsrc/pre_activate/ascend/ir_fusion/momentum_lossscale_fusion.cc +++ b/mindspore/ccsrc/pre_activate/ascend/ir_fusion/momentum_lossscale_fusion.cc @@ -32,7 +32,21 @@ bool CheckValueNodeInputOfMul(const AnfNodePtr &node) { std::vector mul_input_shape = AnfAlgo::GetOutputInferShape(node, 0); return mul_input_shape.empty() || (mul_input_shape.size() == 1 && mul_input_shape[0] == 1); } +void AddInputToOutput(const FuncGraphPtr &func_graph, const CNodePtr &old_cnode, const AnfNodePtr &new_node, + std::vector *new_outputs) { + MS_EXCEPTION_IF_NULL(old_cnode); + MS_EXCEPTION_IF_NULL(new_node); + MS_EXCEPTION_IF_NULL(new_outputs); + auto node_to_output = old_cnode->input(kAccumIndex + 1); + MS_EXCEPTION_IF_NULL(node_to_output); + AbstractBasePtrList abstract_list{old_cnode->abstract(), node_to_output->abstract()}; + auto abstract_tuple = std::make_shared(abstract_list); + new_node->set_abstract(abstract_tuple); + // Create Output + CreateMultipleOutputsOfAnfNode(func_graph, new_node, kFusedMulApplyMomentumOutputNum, new_outputs); +} } // namespace + const BaseRef MomentumLossscaleFusion::DefinePattern() const { VarPtr Xs = std::make_shared(); VarPtr X0 = std::make_shared(); @@ -80,15 +94,10 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph input_names_value[3] = "x1"; input_names_value.emplace_back("x2"); AnfAlgo::SetNodeAttr(kAttrInputNames, MakeValue(input_names_value), new_node); - auto node_to_output = cnode->input(kAccumIndex + 1); - MS_EXCEPTION_IF_NULL(node_to_output); - AbstractBasePtrList abstract_list{node->abstract(), node_to_output->abstract()}; - auto abstract_tuple = std::make_shared(abstract_list); - new_node->set_abstract(abstract_tuple); new_node->set_scope(node->scope()); - // Create Output + // Create Outputs std::vector new_outputs; - CreateMultipleOutputsOfAnfNode(func_graph, new_node, kFusedMulApplyMomentumOutputNum, &new_outputs); + AddInputToOutput(func_graph, cnode, new_node, &new_outputs); if (new_outputs.size() != kFusedMulApplyMomentumOutputNum) { MS_LOG(EXCEPTION) << "Failed to create outputs of " << new_node->DebugString(); }