mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
Modify the code to get correct ragne for symmetric quantization (#10170)
This commit is contained in:
parent
0c517112c4
commit
c29397ad4f
2 changed files with 6 additions and 6 deletions
|
|
@ -860,7 +860,7 @@ class ONNXQuantizer:
|
|||
quantization_params = {}
|
||||
for tensor_name in self.tensors_range.keys():
|
||||
rmin, rmax = self.tensors_range[tensor_name]
|
||||
qmin, qmax = get_qmin_qmax_for_qType(self.input_qType)
|
||||
qmin, qmax = get_qmin_qmax_for_qType(self.input_qType, symmetric=self.is_activation_symmetric)
|
||||
|
||||
quantization_params[tensor_name] = compute_scale_zp(rmin, rmax,
|
||||
qmin, qmax,
|
||||
|
|
|
|||
|
|
@ -176,7 +176,7 @@ def quantize_data(data, qType, symmetric, reduce_range=False):
|
|||
if len(data):
|
||||
rmin = min(data)
|
||||
rmax = max(data)
|
||||
qmin, qmax = get_qmin_qmax_for_qType(qType, reduce_range, for_weight=True)
|
||||
qmin, qmax = get_qmin_qmax_for_qType(qType, reduce_range, symmetric=symmetric)
|
||||
|
||||
zero_point, scale = compute_scale_zp(rmin, rmax, qmin, qmax, symmetric)
|
||||
|
||||
|
|
@ -184,7 +184,7 @@ def quantize_data(data, qType, symmetric, reduce_range=False):
|
|||
|
||||
return rmin, rmax, zero_point, scale, quantized_data
|
||||
|
||||
def get_qmin_qmax_for_qType(qType, reduce_range=False, for_weight=False):
|
||||
def get_qmin_qmax_for_qType(qType, reduce_range=False, symmetric=False):
|
||||
'''
|
||||
Return qmin and qmax, the minimum and maximum value representable by the given qType
|
||||
:parameter qType: onnx.onnx_pb.TensorProto.UINT8 or onnx.onnx_pb.TensorProto.UINT8
|
||||
|
|
@ -193,7 +193,7 @@ def get_qmin_qmax_for_qType(qType, reduce_range=False, for_weight=False):
|
|||
if qType == onnx_proto.TensorProto.UINT8:
|
||||
(qmin, qmax) = (0,127) if reduce_range else (0,255)
|
||||
elif qType == onnx_proto.TensorProto.INT8:
|
||||
if for_weight:
|
||||
if symmetric:
|
||||
(qmin, qmax) = (-64,64) if reduce_range else (-127,127)
|
||||
else:
|
||||
(qmin, qmax) = (-64,64) if reduce_range else (-128,127)
|
||||
|
|
@ -201,13 +201,13 @@ def get_qmin_qmax_for_qType(qType, reduce_range=False, for_weight=False):
|
|||
raise ValueError("Unexpected data type {} requested. Only INT8 and UINT8 are supported.".format(qType))
|
||||
return qmin, qmax
|
||||
|
||||
def get_qrange_for_qType(qType, reduce_range=False, for_weight=False):
|
||||
def get_qrange_for_qType(qType, reduce_range=False, symmetric=False):
|
||||
'''
|
||||
Helper function to get the quantization range for a type.
|
||||
parameter qType: quantization type.
|
||||
return: quantization range.
|
||||
'''
|
||||
qmin, qmax = get_qmin_qmax_for_qType(qType, reduce_range, for_weight=for_weight)
|
||||
qmin, qmax = get_qmin_qmax_for_qType(qType, reduce_range, symmetric=symmetric)
|
||||
return qmax - qmin
|
||||
|
||||
class QuantizedInitializer:
|
||||
|
|
|
|||
Loading…
Reference in a new issue