!32557 Avoid adding duplicate umonad to output of Bprop funcgraph.

Merge pull request !32557 from Margaret_wangrui/bprop_monad
This commit is contained in:
i-robot 2022-04-06 06:53:34 +00:00 committed by Gitee
commit 3b018cecf0
No known key found for this signature in database
GPG Key ID: 173E9B9CA92EEF8F
1 changed files with 23 additions and 1 deletions

View File

@ -471,6 +471,24 @@ MetaFuncGraphPtr KPrim::KMetaFuncGraph(const PrimitivePtr &prim) {
MS_LOG(EXCEPTION) << "Fail to find bprop function for " << prim->name() << ".";
}
bool HasUMonadInput(const CNodePtr &node) {
for (size_t index = node->inputs().size() - 1; index > 0; --index) {
if (HasAbstractUMonad(node->input(index))) {
return true;
}
}
return false;
}
bool HasIOMonadInput(const CNodePtr &node) {
for (size_t index = node->inputs().size() - 1; index > 0; --index) {
if (HasAbstractIOMonad(node->input(index))) {
return true;
}
}
return false;
}
static void AppendMonadOutput(const FuncGraphPtr &bprop_fg, const AnfNodePtr &monad) {
const auto &output = bprop_fg->output();
MS_EXCEPTION_IF_NULL(output);
@ -484,7 +502,11 @@ static void AppendMonadOutput(const FuncGraphPtr &bprop_fg, const AnfNodePtr &mo
}
}
if (output_cnode != nullptr) {
output_cnode->add_input(monad);
bool need_umonad = HasAbstractUMonad(monad) && !HasUMonadInput(output_cnode);
bool need_iomonad = HasAbstractIOMonad(monad) && !HasIOMonadInput(output_cnode);
if (need_umonad || need_iomonad) {
output_cnode->add_input(monad);
}
return;
}
// If output is an empty tuple, create a (make_tuple, monad) as the new output.