From 78ca53e4682b13c2d73a97f6f865c0afa563070f Mon Sep 17 00:00:00 2001 From: mengyuanli Date: Tue, 6 Jul 2021 11:50:02 +0800 Subject: [PATCH] fix bug of control flow --- mindspore/lite/test/config/models_tf.cfg | 3 +- .../optimizer/graph/control_flow_pass.cc | 45 +++++++++---------- .../tools/optimizer/graph/control_flow_pass.h | 6 +-- 3 files changed, 26 insertions(+), 28 deletions(-) diff --git a/mindspore/lite/test/config/models_tf.cfg b/mindspore/lite/test/config/models_tf.cfg index adbead44d8..56d1743b7c 100644 --- a/mindspore/lite/test/config/models_tf.cfg +++ b/mindspore/lite/test/config/models_tf.cfg @@ -91,7 +91,7 @@ ml_video_edit_oneclick_adaptis.pb;3 female_model_step2_int16_noiseout.pb;66 ml_female_model_step6_noiseout.pb;66 ml_male_model_step6_noiseout.pb;66 -#ml_tts_decoder_control_flow.pb;5 need update outputFile +#ml_tts_decoder_control_flow.pb;5 to open ml_tts_decoder.pb;5 ml_tts_encoder_control_flow.pb;4;1:1,22:1:1;;input_dependent ml_tts_vocoder.pb;66 @@ -100,3 +100,4 @@ gts_object_detect_Ics.pb;1;420,630,3;;input_dependent hiai_transformer_encoder.pb;15 decoder_step_nocumsum_v5.pb;13;1:1,512:1,1429,2:1,127:1,127:1,127:1,127,320:1,80:1,512:1,512:1,512:1,512:1,512 hiai_nlu_model_v2.pb;7;1,5:1,6:1,174:1,98:1,5:1,5:1,5 +#ml_audio_kit_encoder_v5.pb;6;1,32:1,32:1,32:1,32:1:1 to open diff --git a/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc b/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc index 5658802a2d..16d99af3cc 100644 --- a/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc @@ -222,8 +222,8 @@ int ControlFlowPass::CreateAfterGraph(const FuncGraphPtr &main_fg, const std::ve } int ControlFlowPass::CreateWhileCondCallNode( - const FuncGraphPtr &fg, const CNodePtr &while_cnode, std::vector *visited_nodes_used_by_after_fg, - CNodePtr *cond_call_cnode, + const FuncGraphPtr &fg, const CNodePtr &while_cnode, const std::vector &visited_nodes_used_by_after_fg, + CNodePtr *cond_call_cnode, std::vector *cond_nodes_used_by_after_partial, std::unordered_map *visited_nodes_and_cond_fg_inputs_replace_pairs) { auto cond_vnode = while_cnode->input(kWhileCondIndex); MS_ASSERT(cond_vnode != nullptr); @@ -245,31 +245,30 @@ int ControlFlowPass::CreateWhileCondCallNode( while_cnode->inputs().end()); auto origin_cond_fg_inputs = cond_fg->get_inputs(); - for (auto it = visited_nodes_used_by_after_fg->begin(); it != visited_nodes_used_by_after_fg->end();) { + for (auto &item : visited_nodes_used_by_after_fg) { bool found = false; - size_t index = -1; + size_t input_index = -1; for (size_t i = kPartialFirstInputSize; i < cond_partial_cnode_inputs.size(); ++i) { - if (cond_partial_cnode_inputs[i] == *it) { + if (cond_partial_cnode_inputs[i] == item) { found = true; - index = i - kPartialFirstInputSize; + input_index = i - kPartialFirstInputSize; break; } } if (found) { - it = visited_nodes_used_by_after_fg->erase(it); - (*visited_nodes_and_cond_fg_inputs_replace_pairs)[*it] = origin_cond_fg_inputs.at(index); + (*visited_nodes_and_cond_fg_inputs_replace_pairs)[item] = origin_cond_fg_inputs.at(input_index); + cond_nodes_used_by_after_partial->push_back(origin_cond_fg_inputs.at(input_index)); continue; } // set after fg inputs to cond_partial_cnode inputs - cond_partial_cnode_inputs.push_back(*it); + cond_partial_cnode_inputs.push_back(item); auto new_parameter = cond_fg->add_parameter(); - new_parameter->set_name((*it)->fullname_with_scope() + "_cond_fg_parameter"); - new_parameter->set_abstract((*it)->abstract()); - (*visited_nodes_and_cond_fg_inputs_replace_pairs)[*it] = new_parameter; - - ++it; + new_parameter->set_name(item->fullname_with_scope() + "_cond_fg_parameter"); + new_parameter->set_abstract(item->abstract()); + (*visited_nodes_and_cond_fg_inputs_replace_pairs)[item] = new_parameter; + cond_nodes_used_by_after_partial->push_back(new_parameter); } auto cond_partial_cnode = fg->NewCNode(cond_partial_cnode_inputs); @@ -349,7 +348,7 @@ int ControlFlowPass::CreateWhileBodyPartialNode(const FuncGraphPtr &cond_fg, con int ControlFlowPass::CreateWhileAfterPartialNode( const FuncGraphPtr &main_fg, const FuncGraphPtr &cond_fg, const std::vector &remain_nodes, - const std::vector &visited_nodes_used_by_after_fg, + const std::vector &cond_nodes_used_by_after_partial, const std::unordered_map &visited_nodes_and_cond_fg_inputs_replace_pairs, CNodePtr *while_cnode, CNodePtr *after_partial_cnode) { // create after_fg @@ -397,20 +396,17 @@ int ControlFlowPass::CreateWhileAfterPartialNode( after_partial_inputs_and_after_fg_inputs_replace_pairs[node] = new_parameter; } - for (auto &pair : after_partial_inputs_and_after_fg_inputs_replace_pairs) { - after_fg->manager()->Replace(pair.first, pair.second); - after_fg->DropNode(pair.first); - } - std::unordered_map visited_nodes_after_fg_replace_pair{}; - for (auto &input : visited_nodes_used_by_after_fg) { + for (auto &input : cond_nodes_used_by_after_partial) { after_partial_cnode_inputs.push_back(visited_nodes_and_cond_fg_inputs_replace_pairs.at(input)); auto new_parameter = after_fg->add_parameter(); new_parameter->set_name(input->fullname_with_scope() + "_after_fg_parameter"); new_parameter->set_abstract(input->abstract()); - visited_nodes_after_fg_replace_pair[input] = new_parameter; + visited_nodes_after_fg_replace_pair[visited_nodes_and_cond_fg_inputs_replace_pairs.at(input)] = new_parameter; } + ReplaceNode(after_fg, visited_nodes_and_cond_fg_inputs_replace_pairs); + ReplaceNode(after_fg, after_partial_inputs_and_after_fg_inputs_replace_pairs); ReplaceNode(after_fg, visited_nodes_after_fg_replace_pair); *after_partial_cnode = cond_fg->NewCNode(after_partial_cnode_inputs); (*after_partial_cnode)->set_fullname_with_scope("CNode_" + after_fg->get_attr("graph_name")->ToString()); @@ -436,8 +432,9 @@ int ControlFlowPass::ProcessWhileOp(const FuncGraphPtr &fg, const std::set visited_nodes_and_cond_fg_inputs_replace_pairs{}; - int ret = CreateWhileCondCallNode(fg, while_cnode, &visited_nodes_used_by_after_fg, &cond_call_cnode, - &visited_nodes_and_cond_fg_inputs_replace_pairs); + std::vector cond_nodes_used_by_after_partial{}; + int ret = CreateWhileCondCallNode(fg, while_cnode, visited_nodes_used_by_after_fg, &cond_call_cnode, + &cond_nodes_used_by_after_partial, &visited_nodes_and_cond_fg_inputs_replace_pairs); if (ret != RET_SUCCESS) { MS_LOG(ERROR) << "while create cond call cnode failed, ret: " << ret; return ret; diff --git a/mindspore/lite/tools/optimizer/graph/control_flow_pass.h b/mindspore/lite/tools/optimizer/graph/control_flow_pass.h index d46bfa7555..beb123ed46 100644 --- a/mindspore/lite/tools/optimizer/graph/control_flow_pass.h +++ b/mindspore/lite/tools/optimizer/graph/control_flow_pass.h @@ -51,13 +51,13 @@ class ControlFlowPass : public Pass { // process while int CreateWhileCondCallNode( - const FuncGraphPtr &fg, const CNodePtr &while_cnode, std::vector *visited_nodes_used_by_after_fg, - CNodePtr *cond_partial_cnode, + const FuncGraphPtr &fg, const CNodePtr &while_cnode, const std::vector &visited_nodes_used_by_after_fg, + CNodePtr *cond_partial_cnode, std::vector *cond_nodes_used_by_after_partial, std::unordered_map *visited_nodes_and_cond_fg_inputs_replace_pairs); int CreateWhileBodyPartialNode(const FuncGraphPtr &cond_fg, const CNodePtr &while_cnode, CNodePtr *body_partial_node); int CreateWhileAfterPartialNode( const FuncGraphPtr &main_fg, const FuncGraphPtr &cond_fg, const std::vector &remain_nodes, - const std::vector &visited_nodes_used_by_after_fg, + const std::vector &cond_nodes_used_by_after_partial, const std::unordered_map &visited_nodes_and_cond_fg_inputs_replace_pairs, CNodePtr *while_cnode, CNodePtr *after_partial_cnode); int ProcessWhileOp(const FuncGraphPtr &fg, const std::set &visited_nodes,