MCPcopy Create free account
hub / github.com/pytorch/pytorch / _test_relu_graph

Method _test_relu_graph

caffe2/python/trt/test_trt.py:85–105  ·  view source on GitHub ↗
(self, X, batch_size, trt_max_batch_size)

Source from the content-addressed store, hash-verified

83 self.opset_version = onnx.defs.onnx_opset_version()
84
85 def _test_relu_graph(self, X, batch_size, trt_max_batch_size):
86 node_def = make_node("Relu", ["X"], ["Y"])
87 Y_c2 = c2.run_node(node_def, {"X": X})
88 graph_def = make_graph(
89 [node_def],
90 name="test",
91 inputs=[make_tensor_value_info("X", onnx.TensorProto.FLOAT, [batch_size, 1, 3, 2])],
92 outputs=[make_tensor_value_info("Y", onnx.TensorProto.FLOAT, [batch_size, 1, 3, 2])])
93 model_def = make_model(graph_def, producer_name='relu-test')
94 op_outputs = [x.name for x in model_def.graph.output]
95 op = convert_onnx_model_to_trt_op(model_def, max_batch_size=trt_max_batch_size)
96 device_option = core.DeviceOption(caffe2_pb2.CUDA, 0)
97 op.device_option.CopyFrom(device_option)
98 Y_trt = None
99 ws = Workspace()
100 with core.DeviceScope(device_option):
101 ws.FeedBlob("X", X)
102 ws.RunOperatorsOnce([op])
103 output_values = [ws.FetchBlob(name) for name in op_outputs]
104 Y_trt = namedtupledict('Outputs', op_outputs)(*output_values)
105 np.testing.assert_almost_equal(Y_c2, Y_trt)
106
107
108 @unittest.skipIf(not workspace.C.use_trt, "No TensortRT support")

Callers 2

Calls 5

WorkspaceClass · 0.90
make_modelFunction · 0.85
namedtupledictFunction · 0.85
run_nodeMethod · 0.45

Tested by

no test coverage detected