[moco] Introduce find_node_byname (#8938)
author박세희/On-Device Lab(SR)/Principal Engineer/삼성전자 <saehie.park@samsung.com>
Thu, 14 Nov 2019 09:04:44 +0000 (18:04 +0900)
committerGitHub Enterprise <noreply-CODE@samsung.com>
Thu, 14 Nov 2019 09:04:44 +0000 (18:04 +0900)
This will intrdocue moco GraphHelper with find_node_byname() that quries node by name.

Signed-off-by: SaeHie Park <saehie.park@samsung.com>
compiler/moco/import/include/moco/GraphHelper.h [new file with mode: 0644]
compiler/moco/import/src/Importer.test.cpp

diff --git a/compiler/moco/import/include/moco/GraphHelper.h b/compiler/moco/import/include/moco/GraphHelper.h
new file mode 100644 (file)
index 0000000..fad62af
--- /dev/null
@@ -0,0 +1,59 @@
+/*
+ * Copyright (c) 2019 Samsung Electronics Co., Ltd. All Rights Reserved
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+#ifndef __MOCO_GRAPH_HELPER_H__
+#define __MOCO_GRAPH_HELPER_H__
+
+#include <moco/IR/TFNode.h>
+
+#include <loco.h>
+
+namespace moco
+{
+
+/**
+ * @brief  find_node_byname() will return a node with type T with given name
+ *         in graph g
+ *
+ * @note   this uses simple linear search, but can speed up with better
+ *         algorithms when needed.
+ */
+template <typename T> T *find_node_byname(loco::Graph *g, const char *name)
+{
+  T *first_node = nullptr;
+  loco::Graph::NodeContext *nodes = g->nodes();
+  uint32_t count = nodes->size();
+
+  for (uint32_t i = 0; i < count; ++i)
+  {
+    auto tfnode = dynamic_cast<TFNode *>(nodes->at(i));
+    if (tfnode != nullptr)
+    {
+      if (tfnode->name() == name)
+      {
+        // if tfnode is NOT type of T then return will be nullptr
+        // this is OK cause the user wanted to get type T but it isn't
+        return dynamic_cast<T *>(tfnode);
+      }
+    }
+  }
+
+  return nullptr;
+}
+
+} // namespace moco
+
+#endif // __MOCO_GRAPH_HELPER_H__
index 6573e8d..2387339 100644 (file)
@@ -15,6 +15,7 @@
  */
 
 #include "moco/Importer.h"
+#include "moco/GraphHelper.h"
 
 #include <moco/IR/Nodes/TFIdentity.h>
 
@@ -210,9 +211,6 @@ TEST(TensorFlowImport, find_node_by_name)
   auto tfidentity = find_first_node_bytype<moco::TFIdentity>(graph.get());
   ASSERT_NE(tfidentity, nullptr);
   ASSERT_NE(tfidentity->input(), nullptr);
-
-// TODO make this test pass
-#if 0
   ASSERT_STREQ(tfidentity->name().c_str(), "output/identity");
 
   auto query_node = moco::find_node_byname<moco::TFConst>(graph.get(), "Foo/w_min");
@@ -222,5 +220,4 @@ TEST(TensorFlowImport, find_node_by_name)
   auto query_node2 = moco::find_node_byname<moco::TFConst>(graph.get(), "Foo/w_max");
   ASSERT_NE(query_node2, nullptr);
   ASSERT_STREQ(query_node2->name().c_str(), "Foo/w_max");
-#endif
 }