| 177 | } |
| 178 | |
| 179 | Status AddShuffleNode(MutableGraphView* graph, const NodeDef& add_before, |
| 180 | const string& buffer_size_node, const string& seed_node, |
| 181 | const string& seed2_node, bool reshuffle_each_iteration) { |
| 182 | NodeDef* add_after = graph->GetNode(add_before.input(0)); |
| 183 | NodeDef new_node; |
| 184 | new_node.set_op(kShuffleDatasetOpName); |
| 185 | graph_utils::SetUniqueGraphNodeName(kShuffleDatasetOpName, graph->graph(), |
| 186 | &new_node); |
| 187 | |
| 188 | new_node.add_input(add_before.input(0)); |
| 189 | new_node.add_input(buffer_size_node); |
| 190 | new_node.add_input(seed_node); |
| 191 | new_node.add_input(seed2_node); |
| 192 | |
| 193 | graph_utils::CopyAttribute("output_shapes", *add_after, &new_node); |
| 194 | graph_utils::CopyAttribute("output_types", *add_after, &new_node); |
| 195 | |
| 196 | AttrValue reshuffle_attr; |
| 197 | reshuffle_attr.set_b(reshuffle_each_iteration); |
| 198 | (*new_node.mutable_attr())["reshuffle_each_iteration"] = reshuffle_attr; |
| 199 | |
| 200 | NodeDef* new_node_graph = graph->AddNode(std::move(new_node)); |
| 201 | |
| 202 | TF_RETURN_IF_ERROR( |
| 203 | graph->UpdateFanouts(add_after->name(), new_node_graph->name())); |
| 204 | return Status::OK(); |
| 205 | } |
| 206 | |
| 207 | Status AddShuffleV2Node(MutableGraphView* graph, const NodeDef& add_before, |
| 208 | const string& buffer_size_node, |
no test coverage detected