Forward inputs at the given indices to outputs and add a control dependency on node.
| 247 | // Forward inputs at the given indices to outputs and add a control dependency |
| 248 | // on node. |
| 249 | bool ConstantFolding::ForwardInputs(NodeDef* node, |
| 250 | absl::Span<const int> inputs_to_forward) { |
| 251 | for (int input_idx : inputs_to_forward) { |
| 252 | if (input_idx < 0 || input_idx >= node->input_size()) { |
| 253 | return false; |
| 254 | } |
| 255 | } |
| 256 | |
| 257 | const std::set<NodeDef*>& tmp = node_map_->GetOutputs(node->name()); |
| 258 | const std::vector<NodeDef*> consumers(tmp.begin(), tmp.end()); |
| 259 | bool updated_graph = false; |
| 260 | for (int input_idx : inputs_to_forward) { |
| 261 | const string& input = node->input(input_idx); |
| 262 | if (IsControlInput(input) && consumers.size() > 1) { |
| 263 | continue; |
| 264 | } |
| 265 | const NodeDef* input_node = node_map_->GetNode(NodeName(input)); |
| 266 | if (input_node == nullptr) { |
| 267 | LOG(ERROR) << "Bad input: " << input; |
| 268 | break; |
| 269 | } |
| 270 | // Update each consumer. |
| 271 | for (NodeDef* consumer : consumers) { |
| 272 | bool add_dep = false; |
| 273 | for (int consumer_input_idx = 0; |
| 274 | consumer_input_idx < consumer->input_size(); ++consumer_input_idx) { |
| 275 | const string& consumer_input = consumer->input(consumer_input_idx); |
| 276 | if (IsControlInput(consumer_input)) { |
| 277 | break; |
| 278 | } |
| 279 | int output_idx; |
| 280 | const string input_node_name = |
| 281 | ParseNodeName(consumer_input, &output_idx); |
| 282 | if (input_node_name == node->name() && output_idx == input_idx) { |
| 283 | consumer->set_input(consumer_input_idx, input); |
| 284 | // We will keep the input from the node through a control |
| 285 | // dependency, so we only need to add the consumer as an output |
| 286 | // for the input node. |
| 287 | node_map_->AddOutput(NodeName(input), consumer->name()); |
| 288 | add_dep = true; |
| 289 | } |
| 290 | } |
| 291 | if (add_dep) { |
| 292 | consumer->add_input(AsControlDependency(node->name())); |
| 293 | updated_graph = true; |
| 294 | } |
| 295 | } |
| 296 | } |
| 297 | |
| 298 | if (updated_graph) { |
| 299 | for (NodeDef* consumer : consumers) { |
| 300 | DedupControlInputs(consumer); |
| 301 | } |
| 302 | } |
| 303 | return updated_graph; |
| 304 | } |
| 305 | |
| 306 | // Puts the given value into the tensor at the given "flat" index. |
nothing calls this directly
no test coverage detected