MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / TEST

Function TEST

tensorflow/core/common_runtime/lower_if_op_test.cc:65–150  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

63}
64
65TEST(LowerIfOpTest, Simple) {
66 std::unique_ptr<Graph> graph(new Graph(OpRegistry::Global()));
67
68 // Add test functions for then and else branch.
69 FunctionDefLibrary f_lib_proto;
70 *(f_lib_proto.add_function()) = test::function::XTimesTwo();
71 *(f_lib_proto.add_function()) = test::function::XTimesFour();
72
73 // Construct simple conditional that switches on `pred` and operates only on
74 // single input `A`.
75 Scope root = Scope::NewRootScope().ExitOnError();
76 TF_ASSERT_OK(root.graph()->AddFunctionLibrary(f_lib_proto));
77 auto a = ops::Placeholder(root.WithOpName("A"), DT_INT32);
78 auto pred = ops::Placeholder(root.WithOpName("pred"), DT_BOOL);
79 Node* written_if;
80 std::vector<NodeBuilder::NodeOut> inputs({NodeBuilder::NodeOut(a.node())});
81 TF_ASSERT_OK(
82 NodeBuilder("if", "If", &root.graph()->flib_def())
83 .Input(pred.node())
84 .Input(inputs)
85 .Attr("then_branch", FuncAttr("XTimesTwo"))
86 .Attr("else_branch", FuncAttr("XTimesFour"))
87 .Attr(LowerFunctionalOpsPass::kLowerUsingSwitchMergeAttr, true)
88 .Attr("Tout", {DT_INT32})
89 .Finalize(root.graph(), &written_if));
90 TF_ASSERT_OK(root.DoShapeInference(written_if));
91 TF_ASSERT_OK(root.ToGraph(graph.get()));
92
93 // The input graph has no switch or merge nodes.
94 int node_called_if_count = 0;
95 for (const auto* op : graph->op_nodes()) {
96 ASSERT_FALSE(op->IsSwitch());
97 ASSERT_FALSE(op->IsMerge());
98 if (op->name() == "if") {
99 ++node_called_if_count;
100 }
101 }
102 ASSERT_EQ(node_called_if_count, 1);
103
104 TF_ASSERT_OK(Rewrite(&graph));
105
106 // Verify the resultant graph has switch and merge nodes, and a node called
107 // `if` (but not If nodes).
108 int switch_count = 0;
109 int merge_count = 0;
110 node_called_if_count = 0;
111 for (const auto* op : graph->op_nodes()) {
112 if (op->IsSwitch()) {
113 ++switch_count;
114 }
115 if (op->IsMerge()) {
116 ++merge_count;
117 }
118 ASSERT_NE(op->type_string(), "If");
119 if (op->name() == "if") {
120 ++node_called_if_count;
121 }
122 }

Callers

nothing calls this directly

Calls 15

XTimesTwoFunction · 0.85
XTimesFourFunction · 0.85
assign_addFunction · 0.85
ExitOnErrorMethod · 0.80
WithOpNameMethod · 0.80
DoShapeInferenceMethod · 0.80
ToGraphMethod · 0.80
op_nodesMethod · 0.80
IsSwitchMethod · 0.80
IsMergeMethod · 0.80
signatureMethod · 0.80
FuncAttrFunction · 0.70

Tested by

no test coverage detected