diff --git a/mindspore/ccsrc/backend/kernel_compiler/tbe/tbe_kernel_build.cc b/mindspore/ccsrc/backend/kernel_compiler/tbe/tbe_kernel_build.cc index b99d1efae0e..cffd182e19e 100644 --- a/mindspore/ccsrc/backend/kernel_compiler/tbe/tbe_kernel_build.cc +++ b/mindspore/ccsrc/backend/kernel_compiler/tbe/tbe_kernel_build.cc @@ -101,6 +101,7 @@ constexpr auto kSOC_VERSION = "SOC_VERSION"; constexpr auto kJIsDynamicShape = "is_dynamic_shape"; constexpr auto kJDynamicIndex = "dynamic_index"; constexpr auto kJSocInfo = "SocInfo"; +constexpr auto kNCHWShapeSize = 4; const auto kPyPath = "/usr/local/Ascend/opp/op_impl/built-in/ai_core/tbe"; @@ -144,7 +145,7 @@ bool TbeKernelJsonCreator::GenTbeSingleKernelJson(const std::shared_ptr &anf_ // !! Note: format: only data node's output use it auto format = AnfAlgo::GetOutputFormat(anf_node, node_out_idx); if (format == kOpFormat_DEFAULT) { - format = ori_shape.size() == 4 ? kOpFormat_NCHW : kOpFormat_ND; + format = ori_shape.size() == kNCHWShapeSize ? kOpFormat_NCHW : kOpFormat_ND; } else if (format == kOpFormat_FRAC_Z) { format = kOpFormat_FRACTAL_Z; } @@ -1365,23 +1366,24 @@ bool TbeKernelBuild::CalOutputSize(const nlohmann::json &fusion_op_list, size_t real_idx = kernel_idx.second; auto full_name = real_node->fullname_with_scope(); for (const auto &op : fusion_op_list) { - if (op[kJName] == full_name) { - auto op_output_desces = op[kJOutputDesc]; - if (output_node != real_node) { - // tuple_get item - auto output_desc = op_output_desces[real_idx]; + if (op[kJName] != full_name) { + continue; + } + auto op_output_desces = op[kJOutputDesc]; + if (output_node != real_node) { + // tuple_get item + auto output_desc = op_output_desces[real_idx]; + if (output_desc[kJShape].empty()) { + MS_LOG(INFO) << "Fusion error: output_desc's shape is empty. real_index " << real_idx; + return false; + } + output_size_list->push_back(GetIOSizeImpl(output_desc)); + } else { + for (const auto &output_desc : op_output_desces) { if (output_desc[kJShape].empty()) { - MS_LOG(INFO) << "Fusion error: output_desc's shape is empty. real_index " << real_idx; - return false; + continue; } output_size_list->push_back(GetIOSizeImpl(output_desc)); - } else { - for (const auto &output_desc : op_output_desces) { - if (output_desc[kJShape].empty()) { - continue; - } - output_size_list->push_back(GetIOSizeImpl(output_desc)); - } } } } diff --git a/mindspore/ccsrc/backend/optimizer/ascend/buffer_fusion/bnupdate_eltwise_fusion_pass.cc b/mindspore/ccsrc/backend/optimizer/ascend/buffer_fusion/bnupdate_eltwise_fusion_pass.cc index 6db793dfd2e..5e5b7f4f9fa 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/buffer_fusion/bnupdate_eltwise_fusion_pass.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/buffer_fusion/bnupdate_eltwise_fusion_pass.cc @@ -27,6 +27,7 @@ namespace mindspore { namespace opt { +constexpr size_t INPUT2 = 2; void BnupdateEltwiseFusionPass::MatchBnupdateDoubleOutputEltwise(const CNodePtr &cnode, const AnfNodePtr &eltwise_input, const session::KernelGraph &kernel_graph, FusedNodeRecord *candidate_fusion) { @@ -48,7 +49,7 @@ void BnupdateEltwiseFusionPass::MatchBnupdateDoubleOutputEltwise(const CNodePtr } auto out_getitem_ptr = out_getitem.first->cast(); MS_EXCEPTION_IF_NULL(out_getitem_ptr); - auto input2 = out_getitem_ptr->input(2); + auto input2 = out_getitem_ptr->input(INPUT2); auto output_idx = GetValue(GetValueNode(input2)); output_used_num[output_idx] = SizeToLong(manager->node_users()[out_getitem.first].size()); } diff --git a/mindspore/ccsrc/backend/optimizer/ascend/format_type/remove_internal_output.cc b/mindspore/ccsrc/backend/optimizer/ascend/format_type/remove_internal_output.cc index fa677181acb..601341f4030 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/format_type/remove_internal_output.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/format_type/remove_internal_output.cc @@ -76,7 +76,7 @@ const AnfNodePtr RemoveInternalOutput::Process(const FuncGraphPtr &func_graph, c } else { auto tuple_getitem = input_node->cast(); MS_EXCEPTION_IF_NULL(tuple_getitem); - int64_t idx = SizeToLong(AnfAlgo::GetTupleGetItemOutIndex(tuple_getitem)); + size_t idx = AnfAlgo::GetTupleGetItemOutIndex(tuple_getitem); AnfNodePtr real_input_node = AnfAlgo::GetTupleGetItemRealInput(tuple_getitem); kernel_graph->ReplaceInternalOutput(node, real_input_node, 0, idx); }