static checking

This commit is contained in:
lz 2021-09-28 11:18:10 +08:00
parent 544ad85a79
commit a4ca032201
5 changed files with 20 additions and 25 deletions

View File

@ -287,6 +287,7 @@ int MindirAdjust::ResetFuncGraph(const FuncGraphPtr &fg, std::set<FuncGraphPtr>
}
bool MindirAdjust::Run(const FuncGraphPtr &func_graph) {
MS_CHECK_TRUE_MSG(func_graph != nullptr, false, "func_graph is nullptr.");
if (this->fmk_type_ != converter::kFmkTypeMs) {
MS_LOG(INFO) << "The framework type of model should be mindir.";
return lite::RET_OK;

View File

@ -33,23 +33,6 @@ constexpr const int kSwitchTruePartialIndex = 2;
constexpr const int kSwitchFalsePartialIndex = 3;
constexpr const int kPartialFgVnodeIndex = 1;
FuncGraphPtr MindIRControlFlowAdjust::GetPartialFg(const CNodePtr &partial_node) {
MS_CHECK_TRUE_MSG(partial_node != nullptr, nullptr, "partial_node is nullptr.");
auto fg_vnode = partial_node->input(kPartialFgVnodeIndex)->cast<ValueNodePtr>();
if (fg_vnode == nullptr) {
MS_LOG(ERROR) << "fg is not right.";
status_ = RET_ERROR;
return nullptr;
}
auto partial_fg = GetValueNode<FuncGraphPtr>(fg_vnode);
if (partial_fg == nullptr) {
MS_LOG(ERROR) << "partial_fg is nullptr.";
status_ = RET_NULL_PTR;
return nullptr;
}
return partial_fg;
}
bool MindIRControlFlowAdjust::HasCallAfter(const FuncGraphPtr &partial_fg) {
MS_CHECK_TRUE_MSG(partial_fg != nullptr, false, "partial_fg is nullptr.");
auto output_node = partial_fg->output();

View File

@ -34,7 +34,6 @@ class MindIRControlFlowAdjust {
bool Run(const FuncGraphPtr &graph);
private:
FuncGraphPtr GetPartialFg(const CNodePtr &partial_node);
std::vector<AnfNodePtr> GetFgOutput(const FuncGraphPtr &fg);
int ModifyFgToCallAfterFg(const FuncGraphPtr &fg, const FuncGraphPtr &after_fg);
bool HasCallAfter(const FuncGraphPtr &partial_fg);

View File

@ -325,9 +325,13 @@ int MoveAttrMapConv2D(const CNodePtr &cnode) {
group = GetValue<int64_t>(dst_prim->GetAttr(ops::kGroup));
}
if (group > 1) {
dst_prim->AddAttr(ops::kIsDepthWise, MakeValue<bool>(true));
auto make_bool_ptr = MakeValue<bool>(true);
MS_CHECK_TRUE_MSG(make_bool_ptr != nullptr, RET_NULL_PTR, "make_bool_ptr is nullptr.");
dst_prim->AddAttr(ops::kIsDepthWise, make_bool_ptr);
}
dst_prim->AddAttr(ops::kGroup, MakeValue(group));
auto make_value_ptr = MakeValue(group);
MS_CHECK_TRUE_MSG(make_value_ptr != nullptr, RET_NULL_PTR, "make_value_ptr is nullptr.");
dst_prim->AddAttr(ops::kGroup, make_value_ptr);
value_node->set_value(dst_prim);
return lite::RET_OK;
}
@ -467,6 +471,7 @@ int MoveAttrMapResize(const CNodePtr &cnode) {
auto dst_prim = std::make_shared<ops::Resize>();
MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr.");
auto size = GetValue<std::vector<int64_t>>(src_prim->GetAttr(ops::kSize));
MS_CHECK_TRUE_MSG(size.size() > 1, RET_ERROR, "out of range.");
dst_prim->set_new_height(size[0]);
dst_prim->set_new_width(size[1]);
if (src_prim->GetAttr(ops::kAlignCorners) != nullptr && GetValue<bool>(src_prim->GetAttr(ops::kAlignCorners))) {
@ -533,6 +538,7 @@ int MoveAttrMapResizeGrad(const CNodePtr &cnode) {
} // namespace
bool PrimitiveAdjust::Run(const FuncGraphPtr &func_graphs) {
MS_ASSERT(func_graphs != nullptr);
if (this->fmk_type_ != converter::kFmkTypeMs) {
MS_LOG(INFO) << "The framework type of model should be mindir.";
return lite::RET_OK;
@ -544,11 +550,17 @@ bool PrimitiveAdjust::Run(const FuncGraphPtr &func_graphs) {
int i = 0;
for (auto func_graph : all_func_graphs) {
func_graph->set_manager(root_func_manager);
func_graph->set_attr("fmk", MakeValue(static_cast<int>(FmkType::kFmkTypeMs)));
auto make_int_ptr = MakeValue(static_cast<int>(FmkType::kFmkTypeMs));
MS_CHECK_TRUE_MSG(make_int_ptr != nullptr, false, "make_int_ptr is nullptr.");
func_graph->set_attr("fmk", make_int_ptr);
if (i == 0) {
func_graph->set_attr("graph_name", MakeValue("main_graph"));
auto make_value_ptr = MakeValue("main_graph");
MS_CHECK_TRUE_MSG(make_value_ptr != nullptr, false, "make_value_ptr is nullptr.");
func_graph->set_attr("graph_name", make_value_ptr);
} else {
func_graph->set_attr("graph_name", MakeValue("subgraph" + std::to_string(i)));
auto make_value_ptr = MakeValue("subgraph" + std::to_string(i));
MS_CHECK_TRUE_MSG(make_value_ptr != nullptr, false, "make_value_ptr is nullptr.");
func_graph->set_attr("graph_name", make_value_ptr);
}
i++;
auto node_list = TopoSort(func_graph->get_return());

View File

@ -64,8 +64,8 @@ int64_t While::get_body_subgraph_index() const {
AbstractBasePtr WhileInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive,
const std::vector<AbstractBasePtr> &input_args) {
MS_CHECK_TRUE_RET(primitive != nullptr, nullptr);
auto While_prim = primitive->cast<PrimWhilePtr>();
MS_CHECK_TRUE_RET(While_prim != nullptr, nullptr);
auto while_prim = primitive->cast<PrimWhilePtr>();
MS_CHECK_TRUE_RET(while_prim != nullptr, nullptr);
AbstractBasePtrList output;
for (int64_t i = 0; i < (int64_t)input_args.size(); i++) {
auto build_shape_ptr = input_args[i]->BuildShape();