| 40 | class RecomputeSubgraphTest : public GrapplerTest {}; |
| 41 | |
| 42 | TEST_F(RecomputeSubgraphTest, SimpleSubgraph) { |
| 43 | tensorflow::Scope s = tensorflow::Scope::NewRootScope(); |
| 44 | |
| 45 | Output a = ops::Variable(s.WithOpName("a"), {2, 3, 4}, DT_FLOAT); |
| 46 | Output b = ops::Identity(s.WithOpName("b"), a); // Recomputed |
| 47 | Output c = ops::Identity(s.WithOpName("c"), b); |
| 48 | Output d = ops::AddN(s.WithOpName("gradients/d"), {c}); |
| 49 | Output e = ops::AddN(s.WithOpName("gradients/e"), {d, b}); |
| 50 | Output f = ops::AddN(s.WithOpName("gradients/f"), {e, a}); |
| 51 | |
| 52 | GrapplerItem item; |
| 53 | TF_CHECK_OK(s.ToGraphDef(&item.graph)); |
| 54 | EXPECT_EQ(6, item.graph.node_size()); |
| 55 | NodeMap pre_transform_node_map(&item.graph); |
| 56 | (*pre_transform_node_map.GetNode("b")->mutable_attr())["_recompute_hint"] |
| 57 | .set_i(0); |
| 58 | |
| 59 | MemoryOptimizer optimizer(RewriterConfig::MANUAL); |
| 60 | GraphDef output; |
| 61 | Status status = optimizer.Optimize(nullptr, item, &output); |
| 62 | |
| 63 | TF_EXPECT_OK(status); |
| 64 | NodeMap post_transform_node_map(&output); |
| 65 | EXPECT_EQ(8, output.node_size()); |
| 66 | NodeDef* transformed_e = post_transform_node_map.GetNode(e.name()); |
| 67 | EXPECT_EQ(2, transformed_e->input_size()); |
| 68 | EXPECT_EQ("gradients/d", transformed_e->input(0)); |
| 69 | EXPECT_EQ("Recomputed/b", transformed_e->input(1)); |
| 70 | NodeDef* recomputed_b = post_transform_node_map.GetNode("Recomputed/b"); |
| 71 | EXPECT_EQ(2, recomputed_b->input_size()); |
| 72 | EXPECT_EQ("a", recomputed_b->input(0)); |
| 73 | EXPECT_EQ("^RecomputeTrigger/b", recomputed_b->input(1)); |
| 74 | NodeDef* recompute_trigger = |
| 75 | post_transform_node_map.GetNode("RecomputeTrigger/b"); |
| 76 | EXPECT_EQ(1, recompute_trigger->input_size()); |
| 77 | EXPECT_EQ("^gradients/d", recompute_trigger->input(0)); |
| 78 | } |
| 79 | |
| 80 | TEST_F(RecomputeSubgraphTest, NoFeedsRecomputed) { |
| 81 | tensorflow::Scope s = tensorflow::Scope::NewRootScope(); |
nothing calls this directly
no test coverage detected