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

Function TEST_F

tensorflow/compiler/tf2tensorrt/convert/convert_graph_test.cc:175–219  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

173};
174
175TEST_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

Callers

nothing calls this directly

Calls 11

WithOpNameMethod · 0.80
nameMethod · 0.65
PlaceholderFunction · 0.50
ShapeClass · 0.50
IdentityFunction · 0.50
AddClass · 0.50
ReshapeFunction · 0.50
InputFunction · 0.50
nodeMethod · 0.45
input_sizeMethod · 0.45
inputMethod · 0.45

Tested by

no test coverage detected