Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 4 additions & 6 deletions py/torch_tensorrt/fx/converters/add.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)

Expand All @@ -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)
6 changes: 2 additions & 4 deletions py/torch_tensorrt/fx/converters/batchnorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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)

Expand Down
14 changes: 0 additions & 14 deletions py/torch_tensorrt/fx/converters/converter_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
3 changes: 1 addition & 2 deletions py/torch_tensorrt/fx/converters/impl/convolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)

Expand Down
7 changes: 3 additions & 4 deletions py/torch_tensorrt/fx/converters/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -39,16 +39,15 @@ 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))
layer.reshape_dims = tuple(input_val.shape[:-1]) + (mod.out_features,)
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)

Expand Down
5 changes: 2 additions & 3 deletions py/torch_tensorrt/fx/converters/mul.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
2 changes: 1 addition & 1 deletion py/torch_tensorrt/fx/converters/quantization.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]

Expand Down
Loading