| 173 | }; |
| 174 | |
| 175 | TEST_F(ConvertAfterShapesTest, DirectlyConnectedEngines) { |
| 176 | // Create the graph. There will be two TRTEngineOps after the conversion, and |
| 177 | // the upstream TRTEngineOp will have two output connections from the same |
| 178 | // node:port inside the op to the downstream TRTEngineOp. Then, if it adds the |
| 179 | // downstream TRTEngineOp first, when adding the upstream op it'll need to |
| 180 | // update the same output connection twice. This test ensures the correctness |
| 181 | // of the conversion under such condition. |
| 182 | Scope s = Scope::NewRootScope(); |
| 183 | auto input = ops::Placeholder(s.WithOpName("input"), DT_FLOAT, |
| 184 | ops::Placeholder::Shape({2, 1})); |
| 185 | // We purposefully choose the name of the root node of each segment, so it'll |
| 186 | // process the segment in the downstream first, then, when it tries to update |
| 187 | // the edge between the two TRTEngineOps, it'll try to add the same edge |
| 188 | // multiple times. |
| 189 | auto segment_root_1 = ops::Identity(s.WithOpName("segment_root_b"), input); |
| 190 | auto add1 = ops::Add(s.WithOpName("add1"), segment_root_1, segment_root_1); |
| 191 | // Add incompatible reshapes that change the batch dimension. |
| 192 | auto incompatible = |
| 193 | ops::Reshape(s.WithOpName("reshape1"), add1, Input({1, 2})); |
| 194 | incompatible = |
| 195 | ops::Reshape(s.WithOpName("reshape2"), incompatible, Input({2, 1})); |
| 196 | |
| 197 | auto add2 = ops::Add(s.WithOpName("add2"), incompatible, add1); |
| 198 | auto segment_root_2 = ops::Identity(s.WithOpName("segment_root_a"), add1); |
| 199 | auto add3 = ops::Add(s.WithOpName("add3"), add2, segment_root_2); |
| 200 | ops::Identity(s.WithOpName("output"), add3); |
| 201 | |
| 202 | GraphDef output_graph_def; |
| 203 | TF_EXPECT_OK(RunConvertAfterShape(s, &output_graph_def)); |
| 204 | |
| 205 | int num_trt_ops = 0; |
| 206 | for (const NodeDef& node : output_graph_def.node()) { |
| 207 | if (node.name() == "TRTEngineOp_1") { |
| 208 | EXPECT_EQ(1, node.input_size()); |
| 209 | EXPECT_EQ("input", node.input(0)); |
| 210 | ++num_trt_ops; |
| 211 | } else if (node.name() == "TRTEngineOp_0") { |
| 212 | EXPECT_EQ(2, node.input_size()); |
| 213 | EXPECT_EQ("TRTEngineOp_1", node.input(0)); |
| 214 | EXPECT_EQ("reshape2", node.input(1)); |
| 215 | ++num_trt_ops; |
| 216 | } |
| 217 | } |
| 218 | EXPECT_EQ(2, num_trt_ops); |
| 219 | } |
| 220 | |
| 221 | } // namespace convert |
| 222 | } // namespace tensorrt |
nothing calls this directly
no test coverage detected