| 592 | } |
| 593 | |
| 594 | Status MutableGraphView::SwapNodeNames(absl::string_view from_node_name, |
| 595 | absl::string_view to_node_name, |
| 596 | bool update_fanouts) { |
| 597 | auto error_status = [from_node_name, to_node_name, |
| 598 | update_fanouts](absl::string_view msg) { |
| 599 | string params = absl::Substitute( |
| 600 | "from_node_name='$0', to_node_name='$1', update_fanouts=$2", |
| 601 | from_node_name, to_node_name, update_fanouts); |
| 602 | return MutationError("SwapNodeNames", params, msg); |
| 603 | }; |
| 604 | |
| 605 | NodeDef* from_node = GetNode(from_node_name); |
| 606 | TF_RETURN_IF_ERROR(CheckNodeExists(from_node_name, from_node, error_status)); |
| 607 | if (from_node_name == to_node_name) { |
| 608 | return Status::OK(); |
| 609 | } |
| 610 | NodeDef* to_node = GetNode(to_node_name); |
| 611 | TF_RETURN_IF_ERROR(CheckNodeExists(to_node_name, to_node, error_status)); |
| 612 | |
| 613 | auto swap_names = [this, from_node, to_node]() { |
| 614 | nodes().erase(from_node->name()); |
| 615 | nodes().erase(to_node->name()); |
| 616 | std::swap(*from_node->mutable_name(), *to_node->mutable_name()); |
| 617 | nodes().emplace(from_node->name(), from_node); |
| 618 | nodes().emplace(to_node->name(), to_node); |
| 619 | }; |
| 620 | |
| 621 | if (update_fanouts) { |
| 622 | SwapFanoutInputs(*this, &fanouts(), &max_regular_output_port(), from_node, |
| 623 | to_node); |
| 624 | swap_names(); |
| 625 | return Status::OK(); |
| 626 | } |
| 627 | |
| 628 | bool from_is_switch = IsSwitch(*from_node); |
| 629 | MutableGraphView::OutputPort to_control(to_node, Graph::kControlSlot); |
| 630 | auto to_control_fanouts = fanouts().find(to_control); |
| 631 | if (from_is_switch && HasFanoutValue(fanouts(), to_control_fanouts)) { |
| 632 | return error_status(SwapNodeNamesSwitchControlErrorMsg(from_node_name)); |
| 633 | } |
| 634 | |
| 635 | bool to_is_switch = IsSwitch(*to_node); |
| 636 | MutableGraphView::OutputPort from_control(from_node, Graph::kControlSlot); |
| 637 | auto from_control_fanouts = fanouts().find(from_control); |
| 638 | if (to_is_switch && HasFanoutValue(fanouts(), from_control_fanouts)) { |
| 639 | return error_status(SwapNodeNamesSwitchControlErrorMsg(to_node_name)); |
| 640 | } |
| 641 | |
| 642 | // Swap node names. |
| 643 | swap_names(); |
| 644 | |
| 645 | // Swap controlling fanouts. |
| 646 | // |
| 647 | // Note: To and from control fanout iterators are still valid as no mutations |
| 648 | // has been performed on fanouts(). |
| 649 | SwapFanoutsMapValues(&fanouts(), from_control, from_control_fanouts, |
| 650 | to_control, to_control_fanouts); |
| 651 | |