| 141 | } |
| 142 | |
| 143 | void StaticShapeCluteringStrategy::Run( |
| 144 | const GraphDef& gdef, |
| 145 | std::vector<std::string>& inputs, |
| 146 | std::vector<std::string>& outputs, |
| 147 | ClusteredGraphInfo* clustered_graph_info) { |
| 148 | std::unordered_map<std::string, bool> dynamic_ops_map = |
| 149 | GetNodesHasDynamicShapeMap(gdef); |
| 150 | |
| 151 | std::unordered_map<std::string, bool> has_control_flow_input = |
| 152 | GetNodesHasControlFlowInputs(gdef); |
| 153 | |
| 154 | std::unordered_map<std::string, const NodeDef*> nodes; |
| 155 | for (const NodeDef& node : gdef.node()) { |
| 156 | nodes[node.name()] = &node; |
| 157 | } |
| 158 | |
| 159 | GraphDef& dynamic_graph_def = clustered_graph_info->tf_subgraph; |
| 160 | GraphDef& static_graph_def = clustered_graph_info->iree_subgraph; |
| 161 | |
| 162 | std::unordered_set<const NodeDef*> static_nodes; |
| 163 | std::unordered_set<const NodeDef*> visited; |
| 164 | std::queue<const NodeDef*> q; |
| 165 | for (auto output : outputs) { |
| 166 | q.push(nodes[output]); |
| 167 | visited.insert(nodes[output]); |
| 168 | } |
| 169 | |
| 170 | std::unordered_set<std::string> black_ops = GetBlackOpsSet(); |
| 171 | while (!q.empty()) { |
| 172 | const NodeDef* curr_node = q.front(); |
| 173 | // 1) no control edge |
| 174 | // 2) no dynamic shape |
| 175 | // 3) no blacklist ops |
| 176 | if (!has_control_flow_input[curr_node->name()] && |
| 177 | !dynamic_ops_map[curr_node->name()] && |
| 178 | black_ops.find(curr_node->op()) == black_ops.end()) { |
| 179 | // Add op into static_graph_def |
| 180 | NodeDef* new_node = static_graph_def.add_node(); |
| 181 | new_node->CopyFrom(*curr_node); |
| 182 | static_nodes.insert(curr_node); |
| 183 | |
| 184 | for (auto in_name : curr_node->input()) { |
| 185 | size_t offset = in_name.find(":"); |
| 186 | in_name = in_name.substr(0, offset); |
| 187 | if (visited.find(nodes[in_name]) == visited.end()) { |
| 188 | q.push(nodes[in_name]); |
| 189 | visited.insert(nodes[in_name]); |
| 190 | } |
| 191 | } |
| 192 | } |
| 193 | |
| 194 | q.pop(); |
| 195 | } |
| 196 | |
| 197 | // TODO: Add version and library ops |
| 198 | |
| 199 | // TODO: Add placeholder here |
| 200 | |