diff --git a/mindspore/ccsrc/pipeline/pipeline.cc b/mindspore/ccsrc/pipeline/pipeline.cc index a20767f77d..63c522ce7b 100644 --- a/mindspore/ccsrc/pipeline/pipeline.cc +++ b/mindspore/ccsrc/pipeline/pipeline.cc @@ -840,6 +840,7 @@ void FinalizeBackend() { void ClearResAtexit() { MS_LOG(DEBUG) << "Pipeline clear all resource"; pynative::ClearPyNativeSession(); + session::ClearPythonParasMap(); device::KernelRuntimeManager::Instance().ClearRuntimeResource(); ad::g_k_prims.clear(); diff --git a/mindspore/ccsrc/session/session_basic.cc b/mindspore/ccsrc/session/session_basic.cc index dd313c6059..8d5ecee79c 100755 --- a/mindspore/ccsrc/session/session_basic.cc +++ b/mindspore/ccsrc/session/session_basic.cc @@ -34,8 +34,30 @@ namespace mindspore { namespace session { +static std::shared_ptr> python_paras_; +void ClearPythonParasMap() { python_paras_ = nullptr; } namespace { const int kSummaryGetItem = 2; + +tensor::TensorPtr GetParamDefaultInputTensor(const AnfNodePtr &node) { + if (node == nullptr) { + return nullptr; + } + auto parameter = node->cast(); + if (parameter == nullptr) { + return nullptr; + } + auto py_param = parameter->default_param(); + if (!py::hasattr(py_param, "default_input")) { + return nullptr; + } + auto py_p_input = py_param.attr("default_input"); + if (!py::hasattr(py_p_input, PYTHON_TENSOR_FLAG)) { + return nullptr; + } + return py_p_input.cast>(); +} + void GetSummaryNodes(const KernelGraph *graph, std::unordered_map> *summary) { MS_LOG(DEBUG) << "Update summary Start"; MS_EXCEPTION_IF_NULL(graph); @@ -195,21 +217,6 @@ ValueNodePtr CreateNewValueNode(const AnfNodePtr &anf, KernelGraph *graph) { return new_value_node; } -ParameterPtr CreateNewParameterFromParameter(const AnfNodePtr &anf, bool valid_input, KernelGraph *graph) { - MS_EXCEPTION_IF_NULL(anf); - if (!anf->isa()) { - MS_LOG(EXCEPTION) << "anf[" << anf->DebugString() << "] is not a parameter"; - } - auto graph_inputs = graph->MutableInputs(); - MS_EXCEPTION_IF_NULL(graph_inputs); - auto valid_inputs = graph->MutableValidInputs(); - MS_EXCEPTION_IF_NULL(valid_inputs); - ParameterPtr new_parameter = graph->NewParameter(anf->cast()); - graph_inputs->push_back(new_parameter); - valid_inputs->push_back(valid_input); - return new_parameter; -} - std::vector CreateParameterFromTuple(const AnfNodePtr &node, bool valid_input, KernelGraph *graph) { MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(graph); @@ -358,6 +365,35 @@ void DumpGraphOutput(const Any &any, size_t recurse_level = 0) { } // namespace GraphId SessionBasic::graph_sum_ = 0; +ParameterPtr SessionBasic::CreateNewParameterFromParameter(const AnfNodePtr &anf, bool valid_input, + KernelGraph *graph) { + MS_EXCEPTION_IF_NULL(anf); + if (!anf->isa()) { + MS_LOG(EXCEPTION) << "anf[" << anf->DebugString() << "] is not a parameter"; + } + + auto m_tensor = GetParamDefaultInputTensor(anf); + auto valid_inputs = graph->MutableValidInputs(); + MS_EXCEPTION_IF_NULL(valid_inputs); + auto graph_inputs = graph->MutableInputs(); + MS_EXCEPTION_IF_NULL(graph_inputs); + ParameterPtr new_parameter = nullptr; + // if parameter's python parameter has been exist a backend parameter, reuse the exist parameter + if (python_paras_ == nullptr) { + python_paras_ = std::make_shared>(); + } + if (python_paras_->find(m_tensor) != python_paras_->end() && GetGraphIdByNode(anf) != kInvalidGraphId) { + new_parameter = (*python_paras_)[m_tensor]; + } else { + new_parameter = graph->NewParameter(anf->cast()); + if (m_tensor != nullptr) { + (*python_paras_)[m_tensor] = new_parameter; + } + } + graph_inputs->push_back(new_parameter); + valid_inputs->push_back(valid_input); + return new_parameter; +} CNodePtr SessionBasic::CreateNewCNode(const CNodePtr &cnode, bool valid_input, KernelGraph *graph, bool *from_other_graph, @@ -391,7 +427,6 @@ CNodePtr SessionBasic::CreateNewCNode(const CNodePtr &cnode, bool valid_input, K } continue; } else if (anf->isa()) { - // if anf is a parameter auto new_parameter = CreateNewParameterFromParameter(anf, valid_input, graph); cnode_inputs.push_back(new_parameter); if (GetGraphIdByNode(anf) == kInvalidGraphId) { diff --git a/mindspore/ccsrc/session/session_basic.h b/mindspore/ccsrc/session/session_basic.h index de443833d6..072226df8f 100755 --- a/mindspore/ccsrc/session/session_basic.h +++ b/mindspore/ccsrc/session/session_basic.h @@ -35,6 +35,7 @@ namespace mindspore { using GraphId = uint32_t; using GraphInfo = std::string; namespace session { +void ClearPythonParasMap(); using CallBackFunc = uint32_t (*)(uint32_t graph_id, const std::map ¶ms_list); using AnyList = std::vector; @@ -106,6 +107,7 @@ class SessionBasic { BaseRef TransformBaseRefListToTuple(const BaseRef &base_ref); // create a new kernel graph and update the graph sum KernelGraphPtr NewKernelGraph(); + ParameterPtr CreateNewParameterFromParameter(const AnfNodePtr &anf, bool valid_input, KernelGraph *graph); std::unordered_map> graphs_; std::unordered_map> run_op_graphs_;