diff --git a/py/torch_tensorrt/fx/converters/add.py b/py/torch_tensorrt/fx/converters/add.py index c60b0313a3..a50b92f27a 100644 --- a/py/torch_tensorrt/fx/converters/add.py +++ b/py/torch_tensorrt/fx/converters/add.py @@ -6,7 +6,7 @@ from ..converter_registry import tensorrt_converter -from .converter_utils import get_dyn_range, mark_as_int8_layer +from .converter_utils import mark_as_int8_layer @tensorrt_converter(operator.add) @@ -43,8 +43,7 @@ def quantized_add(network, target, args, kwargs, layer_name): layer = network.add_elementwise(lhs_val, rhs_val, trt.ElementWiseOperation.SUM) layer.name = layer_name - dyn_range = get_dyn_range(kwargs["scale"], kwargs["zero_point"], torch.quint8) - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) return layer.get_output(0) @@ -60,13 +59,12 @@ def quantized_add_relu(network, submod, args, kwargs, layer_name): layer = network.add_elementwise(lhs_val, rhs_val, trt.ElementWiseOperation.SUM) layer.name = f"{layer_name}_add" - dyn_range = get_dyn_range(kwargs["scale"], kwargs["zero_point"], torch.quint8) - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) layer = network.add_activation( input=layer.get_output(0), type=trt.ActivationType.RELU ) layer.name = f"{layer_name}_relu" - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) return layer.get_output(0) diff --git a/py/torch_tensorrt/fx/converters/batchnorm.py b/py/torch_tensorrt/fx/converters/batchnorm.py index 130991df54..41a91cafd1 100644 --- a/py/torch_tensorrt/fx/converters/batchnorm.py +++ b/py/torch_tensorrt/fx/converters/batchnorm.py @@ -6,7 +6,7 @@ from ..converter_registry import tensorrt_converter -from .converter_utils import get_dyn_range, mark_as_int8_layer, to_numpy +from .converter_utils import mark_as_int8_layer, to_numpy def common_batchnorm(network, mod, input_val, layer_name, is_quantized): @@ -18,9 +18,7 @@ def common_batchnorm(network, mod, input_val, layer_name, is_quantized): layer.name = layer_name if is_quantized: - mark_as_int8_layer( - layer, get_dyn_range(mod.scale, mod.zero_point, torch.quint8) - ) + mark_as_int8_layer(layer) return layer.get_output(0) diff --git a/py/torch_tensorrt/fx/converters/converter_utils.py b/py/torch_tensorrt/fx/converters/converter_utils.py index e6bbff0a5a..1ca01245d2 100644 --- a/py/torch_tensorrt/fx/converters/converter_utils.py +++ b/py/torch_tensorrt/fx/converters/converter_utils.py @@ -693,20 +693,6 @@ def add_reduce_layer( return layer.get_output(0) -def get_dyn_range(scale, zero_point, dtype): - """ - Get the dynamic range of a tensor based on its scale, zero_point and dtype. - """ - if dtype == torch.quint8: - min_val, max_val = 0, 255 - elif dtype == torch.qint8: - min_val, max_val = -128, 127 - else: - raise RuntimeError(f"Unsupported quantized dtype {dtype}") - - return (min_val - zero_point) * scale, (max_val - zero_point) * scale - - def mark_as_int8_layer(layer): """ Set the precision of a layer to int8 as well as the type of its outputs. diff --git a/py/torch_tensorrt/fx/converters/impl/convolution.py b/py/torch_tensorrt/fx/converters/impl/convolution.py index 84071ed2d4..985acb89af 100644 --- a/py/torch_tensorrt/fx/converters/impl/convolution.py +++ b/py/torch_tensorrt/fx/converters/impl/convolution.py @@ -9,7 +9,6 @@ from torch_tensorrt.fx.converters.converter_utils import ( SourceIR, extend_attr_to_tuple, - get_dyn_range, mark_as_int8_layer, set_layer_name, has_dynamic_shape, @@ -127,7 +126,7 @@ def convNd( # Handle quantization cases if scale is not None and zero_point is not None: # Assume the dtype of activation is torch.quint8 - mark_as_int8_layer(conv_layer, get_dyn_range(scale, zero_point, torch.quint8)) + mark_as_int8_layer(conv_layer) result = conv_layer.get_output(0) diff --git a/py/torch_tensorrt/fx/converters/linear.py b/py/torch_tensorrt/fx/converters/linear.py index 1ad6968a15..e6e5461d12 100644 --- a/py/torch_tensorrt/fx/converters/linear.py +++ b/py/torch_tensorrt/fx/converters/linear.py @@ -4,7 +4,7 @@ from ..converter_registry import tensorrt_converter -from .converter_utils import get_dyn_range, mark_as_int8_layer, to_numpy +from .converter_utils import mark_as_int8_layer, to_numpy def common_linear(network, mod, input_val, layer_name, is_quantized): @@ -39,8 +39,7 @@ def common_linear(network, mod, input_val, layer_name, is_quantized): layer.name = f"{layer_name}_linear" if is_quantized: - dyn_range = get_dyn_range(mod.scale, mod.zero_point, torch.quint8) - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) # reshape the output from (*, K, 1, 1) to (*, K) layer = network.add_shuffle(layer.get_output(0)) @@ -48,7 +47,7 @@ def common_linear(network, mod, input_val, layer_name, is_quantized): layer.name = f"{layer_name}_post_shuffle" if is_quantized: - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) return layer.get_output(0) diff --git a/py/torch_tensorrt/fx/converters/mul.py b/py/torch_tensorrt/fx/converters/mul.py index a1d9858ebd..1bfb48de99 100644 --- a/py/torch_tensorrt/fx/converters/mul.py +++ b/py/torch_tensorrt/fx/converters/mul.py @@ -6,7 +6,7 @@ from ..converter_registry import tensorrt_converter -from .converter_utils import get_dyn_range, mark_as_int8_layer +from .converter_utils import mark_as_int8_layer @tensorrt_converter(torch.mul) @@ -42,7 +42,6 @@ def quantized_mul(network, target, args, kwargs, layer_name): layer = network.add_elementwise(lhs_val, rhs_val, trt.ElementWiseOperation.PROD) layer.name = layer_name - dyn_range = get_dyn_range(kwargs["scale"], kwargs["zero_point"], torch.quint8) - mark_as_int8_layer(layer, dyn_range) + mark_as_int8_layer(layer) return layer.get_output(0) diff --git a/py/torch_tensorrt/fx/converters/quantization.py b/py/torch_tensorrt/fx/converters/quantization.py index 968ff955a4..2a0876fcd7 100644 --- a/py/torch_tensorrt/fx/converters/quantization.py +++ b/py/torch_tensorrt/fx/converters/quantization.py @@ -4,7 +4,7 @@ from ..converter_registry import tensorrt_converter -from .converter_utils import get_dyn_range, get_inputs_from_args_and_kwargs +from .converter_utils import get_inputs_from_args_and_kwargs quantize_per_tensor_inputs = ["input", "scale", "zero_point", "dtype"]