| 39 | } // namespace |
| 40 | |
| 41 | Status GraphTopologyView::InitializeFromGraph( |
| 42 | const GraphDef& graph, |
| 43 | const absl::Span<const GraphView::Edge> ephemeral_edges, |
| 44 | bool ignore_control_edges) { |
| 45 | if (graph_ != nullptr) { |
| 46 | return errors::InvalidArgument("GraphTopologyView is already initialized."); |
| 47 | } |
| 48 | |
| 49 | graph_ = &graph; |
| 50 | num_nodes_ = graph.node_size(); |
| 51 | index_to_node_name_.resize(num_nodes_); |
| 52 | node_name_to_index_.rehash(num_nodes_); |
| 53 | fanins_.resize(num_nodes_); |
| 54 | fanouts_.resize(num_nodes_); |
| 55 | |
| 56 | // Build map from name to index and vice versa. |
| 57 | for (int node_idx = 0; node_idx < num_nodes_; ++node_idx) { |
| 58 | const NodeDef& node = graph.node(node_idx); |
| 59 | node_name_to_index_.emplace(node.name(), node_idx); |
| 60 | index_to_node_name_.emplace_back(node.name()); |
| 61 | } |
| 62 | |
| 63 | // 1. Add ephemeral edges to the adjacency lists. |
| 64 | for (const GraphView::Edge& edge : ephemeral_edges) { |
| 65 | const auto src = node_name_to_index_.find(edge.src.node->name()); |
| 66 | const bool valid_src = src != node_name_to_index_.end(); |
| 67 | if (!valid_src) { |
| 68 | const string error_message = |
| 69 | absl::StrCat("Non-existent src node: ", edge.src.node->name()); |
| 70 | if (skip_invalid_edges_) { |
| 71 | VLOG(0) << "Skip error: " << error_message; |
| 72 | } else { |
| 73 | return errors::InvalidArgument(error_message); |
| 74 | } |
| 75 | } |
| 76 | |
| 77 | const auto dst = node_name_to_index_.find(edge.dst.node->name()); |
| 78 | const bool valid_dst = dst != node_name_to_index_.end(); |
| 79 | |
| 80 | if (!valid_dst) { |
| 81 | const string error_message = |
| 82 | absl::StrCat("Non-existent dst node: ", edge.dst.node->name()); |
| 83 | if (skip_invalid_edges_) { |
| 84 | VLOG(0) << "Skip error: " << error_message; |
| 85 | } else { |
| 86 | return errors::InvalidArgument(error_message); |
| 87 | } |
| 88 | } |
| 89 | |
| 90 | if (valid_dst && valid_src) { |
| 91 | const int src_idx = src->second; |
| 92 | const int dst_idx = dst->second; |
| 93 | if (ignore_control_edges && (src_idx < 0 || dst_idx < 0)) { |
| 94 | continue; |
| 95 | } |
| 96 | fanins_[dst_idx].push_back(src_idx); |
| 97 | fanouts_[src_idx].push_back(dst_idx); |
| 98 | } |