diff --git a/py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py b/py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py index bc9a19ab08..c4f9019e59 100644 --- a/py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py +++ b/py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py @@ -4018,6 +4018,7 @@ def aten_ops_arange_start_step( start=args[0], end=args[1], step=args_bounds_check(args, 2, 1), + dtype=kwargs.get("dtype", None), ) diff --git a/py/torch_tensorrt/dynamo/conversion/impl/arange.py b/py/torch_tensorrt/dynamo/conversion/impl/arange.py index f8db7e803f..81fbbae590 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/arange.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/arange.py @@ -1,9 +1,11 @@ from typing import Optional, Union -import numpy as np import tensorrt as trt +import torch from tensorrt import ITensor as TRTTensor +from torch._subclasses.fake_tensor import unset_fake_temporarily from torch.fx.node import Target +from torch_tensorrt import _enums from torch_tensorrt.dynamo.conversion import impl from torch_tensorrt.dynamo.conversion._ConversionContext import ConversionContext from torch_tensorrt.dynamo.conversion.converter_utils import ( @@ -14,63 +16,106 @@ ) +def _sequence_dtype( + dtype: Optional[torch.dtype], + *operands: Union[int, float, torch.Tensor, TRTTensor], +) -> trt.DataType: + """ + Resolve the dtype of the generated sequence. + """ + if dtype is not None: + return _enums.dtype._from(dtype).to(trt.DataType) + + for x in operands: + if isinstance(x, TRTTensor): + if _enums.dtype._from(x.dtype).to(torch.dtype).is_floating_point: + return trt.DataType.FLOAT + + if isinstance(x, torch.Tensor): + if x.dtype.is_floating_point: + return trt.DataType.FLOAT + + if isinstance(x, float): + return trt.DataType.FLOAT + + return trt.DataType.INT64 + + def arange( ctx: ConversionContext, target: Target, source_ir: Optional[SourceIR], name: str, - start: Union[int, TRTTensor], - end: Union[int, TRTTensor], - step: Union[int, TRTTensor], + start: Union[int, float, TRTTensor], + end: Union[int, float, TRTTensor], + step: Union[int, float, TRTTensor], + dtype: Optional[torch.dtype] = None, ) -> TRTTensor: """ Creates a sequence of values (arange) either dynamically or statically, then outputs a TensorRT tensor. If any of (start, end, step) is a TRT tensor, it sets up a dynamic arange - using a Fill layer. Otherwise, it creates a static NumPy array and converts - it into a TensorRT constant tensor. + using a Fill layer. Otherwise, the sequence is computed at build time and + frozen into a TensorRT constant tensor. """ # If any argument is a TRT tensor, use dynamic arange with a Fill layer if any(isinstance(x, TRTTensor) for x in (start, end, step)): - # Convert start, end, step into TRT tensors with appropriate rank - start_rank_0 = get_trt_tensor(ctx, start, name + "_start_rank_0", min_rank=0) + value_dtype = _sequence_dtype(dtype, start, end, step) + start_rank_0 = get_trt_tensor( + ctx, start, name + "_start_rank_0", value_dtype, min_rank=0 + ) # LINSPACE's start input requires rank 0; if the upstream ITensor came in # as rank-1 (e.g. a SymInt materialized by a sym_size op), reshape it. - if isinstance(start_rank_0, TRTTensor) and len(start_rank_0.shape) > 0: + if len(start_rank_0.shape) > 0: squeeze_layer = ctx.net.add_shuffle(start_rank_0) squeeze_layer.reshape_dims = trt.Dims() set_layer_name( squeeze_layer, target, name + "_start_rank_0_squeeze", source_ir ) start_rank_0 = squeeze_layer.get_output(0) - start_rank_1 = get_trt_tensor(ctx, start, name + "_start_rank_1", min_rank=1) - end = get_trt_tensor(ctx, end, name + "_end", min_rank=1) - step = get_trt_tensor(ctx, step, name + "_step", min_rank=1) - # Compute (end - start) / step to determine the output length - shape = impl.elementwise.sub( - ctx, target, source_ir, name + "_sub", end, start_rank_1 + start_rank_1 = get_trt_tensor( + ctx, start, name + "_start_rank_1", value_dtype, min_rank=1 + ) + end = get_trt_tensor(ctx, end, name + "_end", value_dtype, min_rank=1) + step = get_trt_tensor(ctx, step, name + "_step", value_dtype, min_rank=1) + + # The number of elements is ceil((end - start) / step), computed as + # -floor((start - end) / step) so that the whole expression stays in the + # operand dtype and remains a valid TRT shape tensor. Plain truncating + # division would drop the last element whenever the span is not an exact + # multiple of the step. + length = impl.elementwise.sub( + ctx, target, source_ir, name + "_sub", start_rank_1, end + ) + length = impl.elementwise.floor_divide( + ctx, target, source_ir, name + "_floor_div", length, step + ) + length = impl.elementwise.mul( + ctx, target, source_ir, name + "_negate", length, -1 ) - shape = impl.elementwise.trunc_div( - ctx, target, source_ir, name + "_shape", shape, step + length = cast_trt_tensor( + ctx, length, trt.DataType.INT32, name + "_length_casted" ) - shape = cast_trt_tensor(ctx, shape, end.dtype, name + "_shape_casted") # Build a Fill layer in LINSPACE mode fill_layer = ctx.net.add_fill( - shape.shape, trt.FillOperation.LINSPACE, shape.dtype + length.shape, trt.FillOperation.LINSPACE, value_dtype ) - fill_layer.set_input(0, shape) # output length + set_layer_name(fill_layer, target, name + "_arange_fill", source_ir) + fill_layer.set_input(0, length) # output length fill_layer.set_input(1, start_rank_0) # start value fill_layer.set_input(2, step) # step size return fill_layer.get_output(0) else: - # All arguments are static, so use NumPy arange and create a TRT constant - arr = np.arange(start, end, step, dtype=np.int32) - weights = trt.Weights(arr) - const_layer = ctx.net.add_constant(arr.shape, weights) - set_layer_name(const_layer, target, f"{name}_arange_const", source_ir) - return const_layer.get_output(0) + # All arguments are static, so evaluate the sequence eagerly and freeze it + # into the engine as a constant. Letting torch pick the dtype preserves + # PyTorch's promotion rules (float result if any argument is a float). + with unset_fake_temporarily(): + values = torch.arange(start, end, step, dtype=dtype) + if values.dtype == torch.int64: + values = values.to(torch.int32) + return get_trt_tensor(ctx, values, f"{name}_arange_const") diff --git a/tests/py/dynamo/conversion/test_arange_aten.py b/tests/py/dynamo/conversion/test_arange_aten.py index 7705590e2e..e48f14b93e 100644 --- a/tests/py/dynamo/conversion/test_arange_aten.py +++ b/tests/py/dynamo/conversion/test_arange_aten.py @@ -22,6 +22,13 @@ class TestArangeConverter(DispatchTestCase): (-5, -2, 2), (-5, -3, 1), (-2, -5, -1), + (1.2, 5, 1.3), + (1.2, 5.0, 1.3), + (1, 5.0, 1.3), + (-1.2, -5.0, -1.3), + (-5.0, -1.2, 1.3), + (-5, 1.2, 1.3), + (-5.0, 1, 1.3), ] ) def test_arange(self, start, end, step): @@ -36,7 +43,32 @@ def forward(self, x): use_dynamo_tracer=True, ) - def test_arange_dynamic(self): + def test_arange_dynamic_int32(self): + class Arange(nn.Module): + def forward(self, end_tensor): + return torch.ops.aten.arange.start_step(0, end_tensor, 1) + + pyt_input = 7 + inputs = [ + torch_tensorrt.Input( + min_shape=(5,), + opt_shape=(7,), + max_shape=(10,), + dtype=torch.int32, + torch_tensor=torch.tensor(pyt_input, dtype=torch.int32).cuda(), + is_shape_tensor=True, + ) + ] + self.run_test_with_dynamic_shape( + Arange(), + inputs, + use_example_tensors=False, + check_dtype=False, + pyt_inputs=[pyt_input], + use_dynamo_tracer=False, + ) + + def test_arange_dynamic_int64(self): class Arange(nn.Module): def forward(self, end_tensor): return torch.ops.aten.arange.start_step(0, end_tensor, 1)