forked from huawei/mindspore2022
fix:solve input tensor as const for prim bpgraph problem
This commit is contained in:
parent
ac7ce974fb
commit
8b3453e280
|
|
@ -94,6 +94,9 @@ bool CleanAfterOptAPass(const ResourcePtr &res) {
|
|||
}
|
||||
|
||||
FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) {
|
||||
opt::OptPassConfig pynative_eliminate_ = opt::OptPassConfig({
|
||||
irpass.pynative_eliminate_,
|
||||
});
|
||||
opt::irpass::ResolveIRPassLib resolve_irpass;
|
||||
opt::OptPassConfig resolver_prim = opt::OptPassConfig({
|
||||
resolve_irpass.resolver_resolve_and_getattr_,
|
||||
|
|
@ -113,22 +116,21 @@ FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, co
|
|||
irpass.bool_scalar_eliminate,
|
||||
});
|
||||
|
||||
OptPassGroupMap map({{"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 pynative_eliminate_ = opt::OptPassConfig({
|
||||
irpass.pynative_eliminate_,
|
||||
});
|
||||
opt::OptPassConfig switch_simplify_ = opt::OptPassConfig({
|
||||
irpass.switch_simplify_,
|
||||
});
|
||||
|
|
@ -140,11 +142,11 @@ 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_eliminate_", pynative_eliminate_},
|
||||
{"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);
|
||||
|
|
|
|||
|
|
@ -62,16 +62,19 @@ FuncGraphPtr PrimBpropOptimizer::OptimizeBPropFuncGraph(const FuncGraphPtr &bpro
|
|||
}
|
||||
|
||||
PrimitivePtr prim = GetValueNode<PrimitivePtr>(inputs[0]);
|
||||
MS_LOG(WARNING) << "hash of prim " << prim->ToString() << " is:" << prim->hash();
|
||||
|
||||
abstract::AbstractBasePtrList abs_list;
|
||||
ArgsToAbs(op_args, abs_list);
|
||||
ArgsToAbs(prim, op_args, abs_list);
|
||||
|
||||
FuncGraphPtr ret_bprop_fg;
|
||||
PrimBpropOptGraphInfoPtr ret_bprop_info;
|
||||
ECacheQrtRes cache_res = GetOptBpfgFromCache(prim, abs_list, ret_bprop_fg, ret_bprop_info);
|
||||
|
||||
MS_LOG(WARNING) << "cache match result " << cache_res << ", prim: " << prim->ToString();
|
||||
if (cache_res == E_LEVEL_2) {
|
||||
FreeTensorValue(op_args, out, ret_bprop_info);
|
||||
MS_LOG(WARNING)<< "cache level 2 matched, prim: " << prim->ToString();
|
||||
return ret_bprop_fg;
|
||||
}
|
||||
|
||||
|
|
@ -122,6 +125,11 @@ FuncGraphPtr PrimBpropOptimizer::PrimBpropOptStep2(const FuncGraphPtr &bprop_fg,
|
|||
ECacheQrtRes PrimBpropOptimizer::GetOptBpfgFromCache(const PrimitivePtr &prim,
|
||||
const abstract::AbstractBasePtrList &abs_list,
|
||||
FuncGraphPtr &bprop_fg, PrimBpropOptGraphInfoPtr &bprop_info) {
|
||||
auto attrs_ = prim->attrs();
|
||||
for (auto& item : attrs_){
|
||||
MS_LOG(WARNING)<< "attr: " << item.first<<" value:"<<item.second->ToString();
|
||||
}
|
||||
|
||||
auto iter = prim_bprop_cache.find(prim);
|
||||
if (iter == prim_bprop_cache.end()) {
|
||||
return E_NOT_FOUND;
|
||||
|
|
@ -137,37 +145,36 @@ ECacheQrtRes PrimBpropOptimizer::GetOptBpfgFromCache(const PrimitivePtr &prim,
|
|||
return E_LEVEL_2;
|
||||
}
|
||||
|
||||
void PrimBpropOptimizer::ArgsToAbs(const ValuePtrList &op_args, abstract::AbstractBasePtrList &abs_list) {
|
||||
for (auto &item : op_args) {
|
||||
MS_EXCEPTION_IF_NULL(item);
|
||||
auto abs = item->ToAbstract();
|
||||
abs_list.emplace_back(abs);
|
||||
void PrimBpropOptimizer::ArgsToAbs(PrimitivePtr &prim, const ValuePtrList &op_args,
|
||||
abstract::AbstractBasePtrList &abs_list) {
|
||||
auto const_input_index = prim->get_const_input_indexes();
|
||||
bool have_const_input = !const_input_index.empty();
|
||||
bool is_const_prim = prim->is_const_prim();
|
||||
for (size_t i = 0; i < op_args.size(); ++i) {
|
||||
bool is_const_input =
|
||||
have_const_input && std::find(const_input_index.begin(), const_input_index.end(), i) != const_input_index.end();
|
||||
auto &arg_value = op_args[i];
|
||||
auto arg_abs = arg_value->ToAbstract();
|
||||
if (!is_const_prim && !is_const_input) {
|
||||
auto config = abstract::AbstractBase::kBroadenTensorOnly;
|
||||
arg_abs = arg_abs->Broaden(config);
|
||||
MS_LOG(DEBUG) << "Broaden for " << prim->ToString() << " " << config;
|
||||
}
|
||||
abs_list.emplace_back(arg_abs);
|
||||
}
|
||||
}
|
||||
|
||||
void PrimBpropOptimizer::AddOutToAbsList(const ValuePtr &out, abstract::AbstractBasePtrList &abs_list) {
|
||||
if (!out->isa<tensor::Tensor>()) {
|
||||
MS_LOG(EXCEPTION) << "Just suport tensor out now, tuple out need support later.";
|
||||
|
||||
if (!out->isa<tensor::Tensor>() && !out->isa<ValueTuple>()) {
|
||||
MS_LOG(EXCEPTION) << "Out value not Tensor or Tuple, please check the input arguments.";
|
||||
}
|
||||
|
||||
auto tens = out->cast<tensor::TensorPtr>();
|
||||
if (tens->is_parameter()) {
|
||||
abs_list.emplace_back(out->ToAbstract());
|
||||
abs_list.emplace_back(out->ToAbstract());
|
||||
}
|
||||
|
||||
auto dtype = tens->Dtype();
|
||||
if (!IsSubType(dtype, kNumber)) {
|
||||
MS_LOG(EXCEPTION) << "Expect tensor type kNumber but got: " << dtype->ToString() << ".";
|
||||
}
|
||||
auto tensor_shape = tens->shape();
|
||||
auto abs_tensor = std::make_shared<abstract::AbstractTensor>(dtype, tensor_shape);
|
||||
std::string param_name("dout");
|
||||
auto ref_key = std::make_shared<RefKey>(param_name);
|
||||
auto abs_ref_key = ref_key->ToAbstract();
|
||||
auto ref_out = std::make_shared<abstract::AbstractRef>(abs_ref_key, abs_tensor);
|
||||
abs_list.emplace_back(ref_out);
|
||||
abs_list.emplace_back(ref_out);
|
||||
auto out_abs = out->ToAbstract();
|
||||
auto config = abstract::AbstractBase::kBroadenTensorOnly;
|
||||
out_abs = out_abs->Broaden(config);
|
||||
abs_list.emplace_back(out_abs);
|
||||
abs_list.emplace_back(out_abs);
|
||||
}
|
||||
|
||||
FuncGraphPtr OptimizeBPropFuncGraph(const FuncGraphPtr &bprop_fg, const CNodePtr &c_node, const ValuePtrList &op_args,
|
||||
|
|
|
|||
|
|
@ -81,7 +81,7 @@ private:
|
|||
FuncGraphPtr &bprop_fg, PrimBpropOptGraphInfoPtr &bprop_info);
|
||||
|
||||
// converter tensor args to abs value;
|
||||
void ArgsToAbs(const ValuePtrList &op_args, abstract::AbstractBasePtrList &abs_list);
|
||||
void ArgsToAbs(PrimitivePtr &prim, const ValuePtrList &op_args, abstract::AbstractBasePtrList &abs_list);
|
||||
|
||||
// add out && dout to abs list
|
||||
void AddOutToAbsList(const ValuePtr &out, abstract::AbstractBasePtrList &abs_list);
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ class Primitive : public Named {
|
|||
Primitive(const std::string &name, const std::unordered_map<std::string, ValuePtr> &attrs);
|
||||
Primitive(const Primitive &prim);
|
||||
MS_DECLARE_PARENT(Primitive, Named);
|
||||
abstract::AbstractBasePtr ToAbstract();
|
||||
abstract::AbstractBasePtr ToAbstract() override;
|
||||
abstract::AbstractBasePtr ToPrimAbstract(const AnfNodePtr &anf_node);
|
||||
std::string ToString() const override { return name(); }
|
||||
void BeginRecordAddAttr() {
|
||||
|
|
|
|||
Loading…
Reference in New Issue