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

Function TEST

tensorflow/compiler/tf2xla/functionalize_control_flow_test.cc:65–180  ·  view source on GitHub ↗

Graph: x = array_ops.placeholder(dtypes.int32) y = array_ops.placeholder(dtypes.int32) z = control_flow_ops.cond( math_ops.less(y, x), lambda: math_ops.multiply(y, 17), lambda: math_ops.add(x, 23))

Source from the content-addressed store, hash-verified

63// math_ops.less(y, x), lambda: math_ops.multiply(y, 17),
64// lambda: math_ops.add(x, 23))
65TEST(FunctionalizeControlFlow, Conditional) {
66 Graph graph(OpRegistry::Global());
67 {
68 Scope scope = Scope::NewRootScope().ExitOnError();
69
70 auto x = ops::Placeholder(scope.WithOpName("x"), DT_INT32);
71 auto y = ops::Placeholder(scope.WithOpName("y"), DT_INT32);
72 auto less = ops::Less(scope.WithOpName("cond/Less"), y, x);
73 auto switch_1 = ops::Switch(scope.WithOpName("cond/Switch"), less, less);
74
75 auto identity_t =
76 ops::Identity(scope.WithOpName("cond/Identity"), switch_1.output_true);
77 auto seventeen = ops::Const<int32>(
78 scope.WithOpName("cond").WithControlDependencies(identity_t), 17);
79 auto switch_2 = ops::Switch(scope.WithOpName("cond/Switch"), y, less);
80 auto mul = ops::Multiply(scope.WithOpName("cond/Mul"), switch_2.output_true,
81 seventeen);
82
83 auto identity_f =
84 ops::Identity(scope.WithOpName("cond/Identity"), switch_1.output_false);
85 auto twenty_three = ops::Const<int32>(
86 scope.WithOpName("cond").WithControlDependencies(identity_f), 23);
87 auto switch_3 = ops::Switch(scope.WithOpName("cond/Switch"), x, less);
88 auto add = ops::Add(scope.WithOpName("cond/false/add"),
89 switch_3.output_false, twenty_three);
90
91 auto merge = ops::Merge(scope.WithOpName("cond/Merge"),
92 std::initializer_list<Input>{add, mul});
93
94 TF_EXPECT_OK(scope.ToGraph(&graph));
95 }
96
97 FunctionLibraryDefinition library(OpRegistry::Global(), {});
98 GraphDef optimized_graph_def;
99 graph.ToGraphDef(&optimized_graph_def);
100 TF_ASSERT_OK(
101 FunctionalizeControlFlowForGraphDef(&optimized_graph_def, &library));
102 TF_ASSERT_OK(FunctionalizeControlFlow(&graph, &library));
103 GraphDef converted_graph_def;
104 graph.ToGraphDef(&converted_graph_def);
105
106 for (const GraphDef& graph_def : {optimized_graph_def, converted_graph_def}) {
107 string op_name;
108 NameAttrList then_fn;
109 NameAttrList else_fn;
110 TF_EXPECT_OK(FindIfThenAndElse(graph_def, &op_name, &then_fn, &else_fn));
111 InstantiationResultForTest else_result;
112 TF_EXPECT_OK(
113 InstantiateFunctionForTest(else_fn.name(), library, &else_result));
114
115 // Outer graph
116 {
117 Scope scope = Scope::NewRootScope().ExitOnError();
118 auto y = ops::Placeholder(scope.WithOpName("y"), DT_INT32);
119 auto x = ops::Placeholder(scope.WithOpName("x"), DT_INT32);
120 auto less = ops::Less(scope.WithOpName("cond/Less"), y, x);
121 auto if_op = ops::If(scope.WithOpName(op_name), less,
122 std::initializer_list<Input>{less, y, x}, {DT_INT32},

Callers

nothing calls this directly

Calls 15

FunctionalizeControlFlowFunction · 0.85
FindIfThenAndElseFunction · 0.85
IfFunction · 0.85
NextIterationFunction · 0.85
FindWhileCondAndBodyFunction · 0.85
GetNoinlineFunctionDefFunction · 0.85
ExitOnErrorMethod · 0.80
WithOpNameMethod · 0.80
ToGraphMethod · 0.80

Tested by

no test coverage detected