diff --git a/onnxruntime/python/tools/quantization/onnx_model.py b/onnxruntime/python/tools/quantization/onnx_model.py index c156c20293..9a91e2ccee 100644 --- a/onnxruntime/python/tools/quantization/onnx_model.py +++ b/onnxruntime/python/tools/quantization/onnx_model.py @@ -398,6 +398,12 @@ class ONNXModel: return True return False + def is_graph_input(self, tensor_name: str) -> bool: + for input in self.model.graph.input: + if input.name == tensor_name: + return True + return False + # TODO:use OnnxModel.graph_topological_sort(self.model.graph) from transformers.onnx_model # Currently it breaks Openvino/Linux training gpu pipeline so hold off for 1.8 release def topological_sort(self): diff --git a/onnxruntime/python/tools/quantization/qdq_quantizer.py b/onnxruntime/python/tools/quantization/qdq_quantizer.py index 8957b7144f..fb3cc5b3e5 100644 --- a/onnxruntime/python/tools/quantization/qdq_quantizer.py +++ b/onnxruntime/python/tools/quantization/qdq_quantizer.py @@ -201,6 +201,7 @@ class QDQQuantizer(ONNXQuantizer): output_name in self.quantization_params.keys() and len(self.model.input_name_to_nodes()[upstream_output_name]) == 1 and not self.model.is_graph_output(upstream_output_name) + and not self.model.is_graph_input(upstream_output_name) ): self.model.replace_output_of_all_nodes(upstream_output_name, output_name) if upstream_output_name in self.tensors_to_quantize: diff --git a/onnxruntime/test/python/quantization/test_qdq.py b/onnxruntime/test/python/quantization/test_qdq.py index 62965dbe6c..083929673c 100644 --- a/onnxruntime/test/python/quantization/test_qdq.py +++ b/onnxruntime/test/python/quantization/test_qdq.py @@ -6,11 +6,13 @@ # license information. # -------------------------------------------------------------------------- +import tempfile import unittest +from pathlib import Path import numpy as np import onnx -from onnx import TensorProto, helper +from onnx import TensorProto, helper, numpy_helper from op_test_utils import TestDataFeeds, check_model_correctness, check_op_type_count, check_op_type_order from onnxruntime.quantization import QDQQuantizer, QuantFormat, QuantizationMode, QuantType, quantize_static @@ -439,7 +441,74 @@ class TestQDQFormatConvClip(TestQDQFormat): self.verify(True, False) # per_channel:True, is_quant_type_int8:False +def generate_input_initializer(tensor_shape, tensor_dtype, input_name): + """ + Helper function to generate initializers for test inputs + """ + tensor = np.random.normal(0, 0.3, tensor_shape).astype(tensor_dtype) + init = numpy_helper.from_array(tensor, input_name) + return init + + +def construct_relu_conv_model(test_model_path: str) -> None: + """ Create an ONNX model shaped as: + ``` + (input) + | + Relu1 + / \ + Conv1 \ + | \ + Relu2 Conv3 + | | + Conv2 | + \ / + Add + | + (AddOut) + ``` + """ + + input_vi = helper.make_tensor_value_info("input", TensorProto.FLOAT, [1, 3, 1, 3]) + output_vi = helper.make_tensor_value_info("AddOut", TensorProto.FLOAT, [1, 3, 1, 3]) + w1 = generate_input_initializer([3, 3, 1, 1], np.float32, "W1") + b1 = generate_input_initializer([3], np.float32, "B1") + w3 = generate_input_initializer([3, 3, 1, 1], np.float32, "W3") + b3 = generate_input_initializer([3], np.float32, "B3") + w5 = generate_input_initializer([3, 3, 1, 1], np.float32, "W5") + b5 = generate_input_initializer([3], np.float32, "B5") + relu_node_1 = helper.make_node("Relu", ["input"], ["Relu1Out"], name="Relu1") + conv_node_1 = helper.make_node("Conv", ["Relu1Out", "W1", "B1"], ["Conv1Out"], name="Conv1") + relu_node_2 = helper.make_node("Relu", ["Conv1Out"], ["Relu2Out"], name="Relu2") + conv_node_2 = helper.make_node("Conv", ["Relu2Out", "W3", "B3"], ["Conv2Out"], name="Conv2") + conv_node_3 = helper.make_node("Conv", ["Relu1Out", "W5", "B5"], ["Conv3Out"], name="Conv3") + add_node = helper.make_node("Add", ["Conv2Out", "Conv3Out"], ["AddOut"], name="Add") + + graph = helper.make_graph( + [relu_node_1, conv_node_1, relu_node_2, conv_node_2, conv_node_3, add_node], + "test_graph_4", + [input_vi], + [output_vi], + ) + graph.initializer.add().CopyFrom(w1) + graph.initializer.add().CopyFrom(b1) + graph.initializer.add().CopyFrom(w3) + graph.initializer.add().CopyFrom(b3) + graph.initializer.add().CopyFrom(w5) + graph.initializer.add().CopyFrom(b5) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + onnx.save(model, test_model_path) + + class TestQDQFormatConvRelu(TestQDQFormat): + @classmethod + def setUpClass(cls): + cls._tmp_model_dir = tempfile.TemporaryDirectory(prefix="test_qdq_format_conv_relu") + + @classmethod + def tearDownClass(cls): + cls._tmp_model_dir.cleanup() + def construct_model_conv_relu(self, output_model_path, input_shape, weight_shape, output_shape): # (input) # | @@ -482,9 +551,9 @@ class TestQDQFormatConvRelu(TestQDQFormat): def verify(self, per_channel, is_quant_type_int8): np.random.seed(1) - model_fp32_path = "conv_relu_fp32.{}.onnx".format(per_channel) - model_int8_qdq_path = "conv_relu_quant_qdq.{}.onnx".format(per_channel) - model_int8_qop_path = "conv_relu_quant_qop.{}.onnx".format(per_channel) + model_fp32_path = str(Path(self._tmp_model_dir.name) / "conv_relu_fp32.{}.onnx".format(per_channel)) + model_int8_qdq_path = str(Path(self._tmp_model_dir.name) / "conv_relu_quant_qdq.{}.onnx".format(per_channel)) + model_int8_qop_path = str(Path(self._tmp_model_dir.name) / "conv_relu_quant_qop.{}.onnx".format(per_channel)) data_reader = self.input_feeds(1, {"input": [1, 8, 33, 33]}) self.construct_model_conv_relu(model_fp32_path, [1, 8, 33, 33], [16, 8, 3, 3], [1, 16, 31, 31]) quantize_static( @@ -536,6 +605,23 @@ class TestQDQFormatConvRelu(TestQDQFormat): self.verify(False, False) # per_channel:False, is_quant_type_int8:False self.verify(True, False) # per_channel:True, is_quant_type_int8:False + def test_quantize_relu_conv(self): + float_model_path = str(Path(self._tmp_model_dir.name) / "float_relu_convs_model.onnx") + construct_relu_conv_model(float_model_path) + data_reader = self.input_feeds(2, {"input": [1, 3, 1, 3]}) + + qdq_model_path = str(Path(self._tmp_model_dir.name) / "qdq_relu_convs_model.onnx") + quantize_static( + float_model_path, + qdq_model_path, + data_reader, + quant_format=QuantFormat.QDQ, + per_channel=False, + reduce_range=False, + activation_type=QuantType.QInt8, + weight_type=QuantType.QInt8, + ) + if __name__ == "__main__": unittest.main()