Initialize the following graph: w x y z | | | | c1 c2 c3 c4
| 200 | // c1 c2 c3 c4 |
| 201 | // |
| 202 | std::unique_ptr<Graph> InitGraphForPruning() { |
| 203 | GraphDefBuilder builder(GraphDefBuilder::kFailImmediately); |
| 204 | const string dev0 = "/job:localhost/replica:0/task:0/device:CPU:0"; |
| 205 | Node* w = ops::SourceOp("TestParams", |
| 206 | builder.opts().WithName("w").WithDevice(dev0)); |
| 207 | Node* x = ops::SourceOp("TestParams", |
| 208 | builder.opts().WithName("x").WithDevice(dev0)); |
| 209 | Node* y = ops::SourceOp("TestParams", |
| 210 | builder.opts().WithName("y").WithDevice(dev0)); |
| 211 | Node* z = ops::SourceOp("TestParams", |
| 212 | builder.opts().WithName("z").WithDevice(dev0)); |
| 213 | CollectiveReduceNode(&builder, w, "c1", dev0, 1); |
| 214 | CollectiveReduceNode(&builder, x, "c2", dev0, 2); |
| 215 | CollectiveReduceNode(&builder, y, "c3", dev0, 3); |
| 216 | CollectiveReduceNode(&builder, z, "c4", dev0, 4); |
| 217 | |
| 218 | std::unique_ptr<Graph> graph = absl::make_unique<Graph>(OpRegistry::Global()); |
| 219 | Status s = GraphDefBuilderToGraph(builder, graph.get()); |
| 220 | if (!s.ok()) { |
| 221 | LOG(FATAL) << "Error building graph " << s; |
| 222 | } |
| 223 | return graph; |
| 224 | } |
| 225 | |
| 226 | // Tests that in the graph created by `InitGraphForPruning`, we only add c4 -> |
| 227 | // c3, c3 -> c2, c2 -> c1, and other edges are pruned away. |
no test coverage detected