| 1234 | } |
| 1235 | |
| 1236 | Status MutableGraphView::UpdateFanin(absl::string_view node_name, |
| 1237 | const TensorId& from_fanin, |
| 1238 | const TensorId& to_fanin) { |
| 1239 | auto error_status = [node_name, from_fanin, to_fanin](absl::string_view msg) { |
| 1240 | string params = |
| 1241 | absl::Substitute("node_name='$0', from_fanin='$1', to_fanin='$2'", |
| 1242 | node_name, from_fanin.ToString(), to_fanin.ToString()); |
| 1243 | return MutationError("UpdateFanin", params, msg); |
| 1244 | }; |
| 1245 | |
| 1246 | TF_RETURN_IF_ERROR(CheckFaninIsValid(from_fanin, error_status)); |
| 1247 | TF_RETURN_IF_ERROR(CheckFaninIsValid(to_fanin, error_status)); |
| 1248 | NodeDef* node = GetNode(node_name); |
| 1249 | TF_RETURN_IF_ERROR(CheckNodeExists(node_name, node, error_status)); |
| 1250 | NodeDef* from_fanin_node = GetNode(from_fanin.node()); |
| 1251 | TF_RETURN_IF_ERROR( |
| 1252 | CheckNodeExists(from_fanin.node(), from_fanin_node, error_status)); |
| 1253 | NodeDef* to_fanin_node = GetNode(to_fanin.node()); |
| 1254 | TF_RETURN_IF_ERROR( |
| 1255 | CheckNodeExists(to_fanin.node(), to_fanin_node, error_status)); |
| 1256 | |
| 1257 | // When replacing a non control dependency fanin with a control dependency, or |
| 1258 | // vice versa, remove and add, so ports can be updated properly in fanout(s). |
| 1259 | bool to_fanin_is_control = IsTensorIdControlling(to_fanin); |
| 1260 | if (to_fanin_is_control && IsSwitch(*to_fanin_node)) { |
| 1261 | // Can't add Switch node as a control dependency. |
| 1262 | return error_status( |
| 1263 | absl::Substitute("can't update to fanin '$0' as it will become a " |
| 1264 | "Switch control dependency", |
| 1265 | to_fanin.ToString())); |
| 1266 | } |
| 1267 | if (node_name == from_fanin.node() || node_name == to_fanin.node()) { |
| 1268 | return error_status("can't update fanin to or from self"); |
| 1269 | } |
| 1270 | |
| 1271 | if (from_fanin == to_fanin) { |
| 1272 | return Status::OK(); |
| 1273 | } |
| 1274 | |
| 1275 | bool from_fanin_is_control = IsTensorIdControlling(from_fanin); |
| 1276 | if (from_fanin_is_control || to_fanin_is_control) { |
| 1277 | bool modified = false; |
| 1278 | if (from_fanin_is_control) { |
| 1279 | modified |= RemoveControllingFaninInternal(node, from_fanin_node); |
| 1280 | } else { |
| 1281 | modified |= RemoveRegularFaninInternal( |
| 1282 | node, {from_fanin_node, from_fanin.index()}); |
| 1283 | } |
| 1284 | if (modified) { |
| 1285 | AddFaninInternal(node, {to_fanin_node, to_fanin.index()}); |
| 1286 | } |
| 1287 | return Status::OK(); |
| 1288 | } |
| 1289 | |
| 1290 | // In place mutation of regular fanins, requires no shifting of ports. |
| 1291 | string to_fanin_string = TensorIdToString(to_fanin); |
| 1292 | const int num_regular_fanins = |
| 1293 | NumFanins(*node, /*include_controlling_nodes=*/false); |