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))
| 63 | // math_ops.less(y, x), lambda: math_ops.multiply(y, 17), |
| 64 | // lambda: math_ops.add(x, 23)) |
| 65 | TEST(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}, |
nothing calls this directly
no test coverage detected