| 378 | } |
| 379 | |
| 380 | NodeDef* GetTailOfChain(const NodeDef& source, const NodeMap& node_map, |
| 381 | bool follow_control_input, |
| 382 | const std::function<bool(const NodeDef&)>& pred_fn) { |
| 383 | const NodeDef* current = &source; |
| 384 | const NodeDef* next = current; |
| 385 | while (next == &source || (next != nullptr && pred_fn(*next))) { |
| 386 | current = next; |
| 387 | if (current->input_size() == 0 || |
| 388 | (!follow_control_input && IsControlInput(current->input(0)))) { |
| 389 | break; |
| 390 | } |
| 391 | next = node_map.GetNode(current->input(0)); |
| 392 | if (next == nullptr) { |
| 393 | LOG(ERROR) << "Node not found: " << current->input(0); |
| 394 | } |
| 395 | } |
| 396 | return const_cast<NodeDef*>(current); |
| 397 | } |
| 398 | |
| 399 | // Every permutation is a product of one or more cycles. Iterate over the cycles |
| 400 | // in the permutation, and convert each of those into a product of |