mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
fix qdq relu removal bug (#12542)
Fix minor bug in qdq quantization tool Motivation and Context Relu node is removed in qdq quantization tool if it can be merged to its input node. When performing the removal, we forgot to check whether the input is actually the graph input
This commit is contained in:
parent
c10704a501
commit
b2382dc43a
3 changed files with 97 additions and 4 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Reference in a new issue