From 76bc0734d51ac53e4dc94e4a77472d7c9fd6882b Mon Sep 17 00:00:00 2001 From: lvliang Date: Wed, 24 Feb 2021 11:38:47 +0800 Subject: [PATCH] modify_codes_to_call_kpynativecellend_in_GradNetInner --- .../ccsrc/frontend/optimizer/ad/kpynative.cc | 17 ++++---- .../ccsrc/frontend/optimizer/ad/kpynative.h | 2 - mindspore/ccsrc/pipeline/jit/pass.cc | 40 +++++++++++++------ mindspore/ccsrc/pipeline/jit/pass.h | 1 + .../pipeline/jit/prim_bprop_optimizer.cc | 22 +++++++--- .../ccsrc/pipeline/jit/prim_bprop_optimizer.h | 17 ++++---- mindspore/ccsrc/pipeline/pynative/base.h | 1 - 7 files changed, 64 insertions(+), 36 deletions(-) diff --git a/mindspore/ccsrc/frontend/optimizer/ad/kpynative.cc b/mindspore/ccsrc/frontend/optimizer/ad/kpynative.cc index abe3493479e..73ac8c4bf3b 100644 --- a/mindspore/ccsrc/frontend/optimizer/ad/kpynative.cc +++ b/mindspore/ccsrc/frontend/optimizer/ad/kpynative.cc @@ -20,6 +20,7 @@ #include #include #include "ir/anf.h" +#include "pipeline/jit/prim_bprop_optimizer.h" #include "frontend/optimizer/ad/adjoint.h" #include "frontend/optimizer/ad/dfunctor.h" #include "frontend/optimizer/ad/kpynative.h" @@ -28,6 +29,7 @@ #include "utils/primitive_utils.h" #include "utils/ms_context.h" #include "utils/info.h" +#include "debug/anf_ir_dump.h" #include "debug/trace.h" namespace mindspore { @@ -47,7 +49,6 @@ class KPynativeCellImpl : public KPynativeCell { bool KPynativeWithBProp(const CNodePtr &c_node, const ValuePtrList &op_args, const ValuePtr &out, const FuncGraphPtr &bprop_fg); FuncGraphPtr Finish(const AnfNodePtrList &weights, bool grad_inputs, bool grad_weights); - FuncGraphPtr bg() { return tape_; } private: FuncGraphPtr tape_; @@ -66,11 +67,6 @@ KPynativeCellPtr GradPynativeCellBegin(const AnfNodePtrList &cell_inputs) { return std::make_shared(cell_inputs); } -FuncGraphPtr GetPynativeBg(const KPynativeCellPtr &k_cell) { - auto k_cell_impl = std::dynamic_pointer_cast(k_cell); - return k_cell_impl->bg(); -} - FuncGraphPtr GradPynativeCellEnd(const KPynativeCellPtr &k_cell, const AnfNodePtrList &weights, bool grad_inputs, bool grad_weights) { auto k_cell_impl = std::dynamic_pointer_cast(k_cell); @@ -129,6 +125,11 @@ FuncGraphPtr KPynativeCellImpl::Finish(const AnfNodePtrList &weights, bool grad_ } tr.Commit(); + // Do inline opt for final bprop graph + DumpIR("before_final_inline.ir", tape_); + tape_ = pipeline::PrimBpropOptimizer::GetPrimBpropOptimizerInst().BpropGraphInlineOpt(tape_); + DumpIR("after_final_inline.ir", tape_); + return tape_; } @@ -172,7 +173,9 @@ bool KPynativeCellImpl::KPynativeWithBProp(const CNodePtr &c_node, const ValuePt FuncGraphPtr OptimizeBPropFuncGraph(const FuncGraphPtr &bprop_fg, const CNodePtr &c_node, const ValuePtrList &op_args, const ValuePtr &out) { - return bprop_fg; + auto optimized_bprop_fg = + pipeline::PrimBpropOptimizer::GetPrimBpropOptimizerInst().OptimizeBPropFuncGraph(bprop_fg, c_node, op_args, out); + return optimized_bprop_fg; } bool KPynativeCellImpl::BackPropagate(const CNodePtr &cnode_primal, const CNodePtr &bprop_app) { diff --git a/mindspore/ccsrc/frontend/optimizer/ad/kpynative.h b/mindspore/ccsrc/frontend/optimizer/ad/kpynative.h index 8ab935d7f6e..30516a9fcd8 100644 --- a/mindspore/ccsrc/frontend/optimizer/ad/kpynative.h +++ b/mindspore/ccsrc/frontend/optimizer/ad/kpynative.h @@ -70,8 +70,6 @@ bool GradPynativeOp(const KPynativeCellPtr &k_cell, const CNodePtr &c_node, cons // Should have prototype: (sens_input1, sens_input2, ...) bprop_fg(input1, input2, ..., out, dout) bool GradPynativeWithBProp(const KPynativeCellPtr &k_cell, const CNodePtr &c_node, const ValuePtrList &op_args, const ValuePtr &out, const FuncGraphPtr &bprop_fg); - -FuncGraphPtr GetPynativeBg(const KPynativeCellPtr &k_cell); } // namespace ad } // namespace mindspore diff --git a/mindspore/ccsrc/pipeline/jit/pass.cc b/mindspore/ccsrc/pipeline/jit/pass.cc index 7363650864f..a306899b1f0 100644 --- a/mindspore/ccsrc/pipeline/jit/pass.cc +++ b/mindspore/ccsrc/pipeline/jit/pass.cc @@ -116,20 +116,20 @@ FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, co irpass.bool_scalar_eliminate, }); - OptPassGroupMap map({ - {"ad_eliminate_", pynative_eliminate_}, - {"ad_resolver_prim", resolver_prim}, - {"ad_inline_", inline_}, - {"bool_scalar_eliminate", bool_scalar_eliminate}, - {"ad_switch_simplify_", switch_simplify_}}); + OptPassGroupMap map({{"ad_eliminate_", pynative_eliminate_}, + {"ad_resolver_prim", resolver_prim}, + {"ad_inline_", inline_}, + {"bool_scalar_eliminate", bool_scalar_eliminate}, + {"ad_switch_simplify_", switch_simplify_}}); auto prim_bprop_opt_step_1 = opt::Optimizer::MakeOptimizer("prim_bprop_opt_step_1", res, map); FuncGraphPtr func_graph = res->func_graph(); - WITH(MsProfile::GetProfile()->Step("prim_bprop_opt_step_1")) [&prim_bprop_opt_step_1, &func_graph]() { + WITH(MsProfile::GetProfile()->Step("prim_bprop_opt_step_1"))[&prim_bprop_opt_step_1, &func_graph]() { func_graph = prim_bprop_opt_step_1->step(func_graph, true); }; return func_graph; } + FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) { opt::OptPassConfig switch_simplify_ = opt::OptPassConfig({ irpass.switch_simplify_, @@ -142,12 +142,10 @@ FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, co auto re_auto_monadwrapper = [](const FuncGraphPtr &root, const opt::OptimizerPtr &) -> bool { return ReAutoMonad(root); }; - OptPassGroupMap map({ - {"ad_renormalize", opt::OptPassConfig::Renormalize()}, - {"ad_inline_", inline_}, - {"ad_switch_simplify_", switch_simplify_}, - {"auto_monad_grad", opt::OptPassConfig(re_auto_monadwrapper)} - }); + OptPassGroupMap map({{"ad_renormalize", opt::OptPassConfig::Renormalize()}, + {"ad_inline_", inline_}, + {"ad_switch_simplify_", switch_simplify_}, + {"auto_monad_grad", opt::OptPassConfig(re_auto_monadwrapper)}}); auto prim_bprop_opt_step_2 = opt::Optimizer::MakeOptimizer("prim_bprop_opt_step_2", res, map); FuncGraphPtr func_graph = res->func_graph(); @@ -157,6 +155,22 @@ FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, co return func_graph; } +FuncGraphPtr BpropGraphInlineOptPass(const ResourcePtr &res) { + opt::irpass::OptimizeIRPassLib irpass; + opt::OptPassConfig inline_ = opt::OptPassConfig({ + irpass.inline_, + }); + + OptPassGroupMap map({{"ad_inline_", inline_}}); + + auto bprop_graph_inline_opt = opt::Optimizer::MakeOptimizer("bprop_graph_inline_opt", res, map); + FuncGraphPtr func_graph = res->func_graph(); + WITH(MsProfile::GetProfile()->Step("bprop_graph_inline_opt"))[&bprop_graph_inline_opt, &func_graph]() { + func_graph = bprop_graph_inline_opt->step(func_graph, true); + }; + return func_graph; +} + namespace { bool ReAutoMonadWrapper(const FuncGraphPtr &root, const opt::OptimizerPtr &) { return ReAutoMonad(root); } diff --git a/mindspore/ccsrc/pipeline/jit/pass.h b/mindspore/ccsrc/pipeline/jit/pass.h index eaf53b49af6..9dd9cb15433 100644 --- a/mindspore/ccsrc/pipeline/jit/pass.h +++ b/mindspore/ccsrc/pipeline/jit/pass.h @@ -48,6 +48,7 @@ void ReclaimOptimizer(); bool PynativeOptPass(const ResourcePtr &res); FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res); FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res); +FuncGraphPtr BpropGraphInlineOptPass(const ResourcePtr &res); } // namespace pipeline } // namespace mindspore diff --git a/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.cc b/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.cc index 0a7a659cfd4..9c9d5687596 100644 --- a/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.cc +++ b/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.cc @@ -31,8 +31,7 @@ PrimBpropOptimizer::PrimBpropOptimizer() { prim_bprop_opt_manage = prim_bprop_opt_res->manager(); } -PrimBpropOptimizer::~PrimBpropOptimizer() { -} +PrimBpropOptimizer::~PrimBpropOptimizer() {} void PrimBpropOptimizer::Clear() { prim_bprop_cache.clear(); @@ -122,6 +121,20 @@ FuncGraphPtr PrimBpropOptimizer::PrimBpropOptStep2(const FuncGraphPtr &bprop_fg, return opt_bprop_fg; } +FuncGraphPtr PrimBpropOptimizer::BpropGraphInlineOpt(const FuncGraphPtr &bprop_fg) { + MS_EXCEPTION_IF_NULL(bprop_fg); + auto bprop_graph_opt_res = std::make_shared(); + auto bprop_graph_opt_manage = bprop_graph_opt_res->manager(); + bprop_graph_opt_res->set_func_graph(bprop_fg); + bprop_graph_opt_manage->AddFuncGraph(bprop_fg); + auto after_inline_bg = BpropGraphInlineOptPass(bprop_graph_opt_res); + // Clear resource + bprop_graph_opt_manage->Clear(); + bprop_graph_opt_res->Clean(); + + return after_inline_bg; +} + ECacheQrtRes PrimBpropOptimizer::GetOptBpfgFromCache(const PrimitivePtr &prim, const abstract::AbstractBasePtrList &abs_list, FuncGraphPtr &bprop_fg, PrimBpropOptGraphInfoPtr &bprop_info) { @@ -164,9 +177,8 @@ void PrimBpropOptimizer::ArgsToAbs(PrimitivePtr &prim, const ValuePtrList &op_ar } } -abstract::AbstractBasePtrList PrimBpropOptimizer::AddOutToAbsList( - const ValuePtr &out, const abstract::AbstractBasePtrList &abs_list) { - +abstract::AbstractBasePtrList PrimBpropOptimizer::AddOutToAbsList(const ValuePtr &out, + const abstract::AbstractBasePtrList &abs_list) { if (!out->isa() && !out->isa()) { MS_LOG(EXCEPTION) << "Out value not Tensor or Tuple, please check the input arguments."; } diff --git a/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.h b/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.h index 039649e119b..cebaaf77f24 100644 --- a/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.h +++ b/mindspore/ccsrc/pipeline/jit/prim_bprop_optimizer.h @@ -31,7 +31,7 @@ using PrimBpropOptGraphInfoPtr = std::shared_ptr; using PrimBpropCache = std::unordered_map; using AbstractListMap = std::unordered_map; + abstract::AbstractBasePtrListHasher, abstract::AbstractBasePtrListEqual>; struct PrimitiveTotalEqual { bool operator()(PrimitivePtr const &t1, PrimitivePtr const &t2) const { @@ -41,9 +41,7 @@ struct PrimitiveTotalEqual { } }; -enum ECacheQrtRes { - E_NOT_FOUND, E_LEVEL_1, E_LEVEL_2 -}; +enum ECacheQrtRes { E_NOT_FOUND, E_LEVEL_1, E_LEVEL_2 }; struct PrimBpropOptGraphInfo { // the opt funcgraph without infer, level1 cache @@ -56,7 +54,7 @@ struct PrimBpropOptGraphInfo { }; class PrimBpropOptimizer { -public: + public: ~PrimBpropOptimizer(); void Clear(); @@ -71,10 +69,13 @@ public: FuncGraphPtr OptimizeBPropFuncGraph(const FuncGraphPtr &bprop_fg, const CNodePtr &c_node, const ValuePtrList &op_args, const ValuePtr &out); + // do inline opt for final bprop graph + FuncGraphPtr BpropGraphInlineOpt(const FuncGraphPtr &bprop_fg); + // need ? how to shrink ? // void CacheShrink(); -private: + private: PrimBpropOptimizer(); ECacheQrtRes GetOptBpfgFromCache(const PrimitivePtr &prim, const abstract::AbstractBasePtrList &abs_list, @@ -87,7 +88,7 @@ private: abstract::AbstractBasePtrList AddOutToAbsList(const ValuePtr &out, const abstract::AbstractBasePtrList &abs_list); // TODO: how To? - void FreeTensorValue(const ValuePtrList &op_args, const ValuePtr &out, PrimBpropOptGraphInfoPtr &bprop_info) {}; + void FreeTensorValue(const ValuePtrList &op_args, const ValuePtr &out, PrimBpropOptGraphInfoPtr &bprop_info){}; // do opt without input info, no infer FuncGraphPtr PrimBpropOptStep1(const FuncGraphPtr &bprop_fg); @@ -97,7 +98,7 @@ private: void BindAbsToParameters(const FuncGraphPtr &bprop_fg, abstract::AbstractBasePtrList &abs_list_input); -private: + private: FuncGraphManagerPtr prim_bprop_opt_manage; ResourcePtr prim_bprop_opt_res; // cache optimized bprop graph diff --git a/mindspore/ccsrc/pipeline/pynative/base.h b/mindspore/ccsrc/pipeline/pynative/base.h index f858ceb80c4..52cc91a4077 100644 --- a/mindspore/ccsrc/pipeline/pynative/base.h +++ b/mindspore/ccsrc/pipeline/pynative/base.h @@ -51,7 +51,6 @@ enum RunOpArgsEnum { PY_PRIM = 0, PY_NAME, PY_INPUTS, PY_ARGS_NUM }; struct OpExecInfo { std::string op_name; - std::string op_index; PrimitivePyPtr py_primitive; AbstractBasePtr abstract;