forked from huawei/mindspore2022
modify_codes_to_call_kpynativecellend_in_GradNetInner
This commit is contained in:
parent
e9e00450d5
commit
76bc0734d5
|
|
@ -20,6 +20,7 @@
|
|||
#include <string>
|
||||
#include <utility>
|
||||
#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<KPynativeCellImpl>(cell_inputs);
|
||||
}
|
||||
|
||||
FuncGraphPtr GetPynativeBg(const KPynativeCellPtr &k_cell) {
|
||||
auto k_cell_impl = std::dynamic_pointer_cast<KPynativeCellImpl>(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<KPynativeCellImpl>(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) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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); }
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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<pipeline::Resource>();
|
||||
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<tensor::Tensor>() && !out->isa<ValueTuple>()) {
|
||||
MS_LOG(EXCEPTION) << "Out value not Tensor or Tuple, please check the input arguments.";
|
||||
}
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ using PrimBpropOptGraphInfoPtr = std::shared_ptr<PrimBpropOptGraphInfo>;
|
|||
using PrimBpropCache = std::unordered_map<PrimitivePtr, PrimBpropOptGraphInfoPtr, PrimitiveHasher, PrimitiveTotalEqual>;
|
||||
|
||||
using AbstractListMap = std::unordered_map<abstract::AbstractBasePtrList, FuncGraphPtr,
|
||||
abstract::AbstractBasePtrListHasher, abstract::AbstractBasePtrListEqual>;
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue