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
1 change: 1 addition & 0 deletions py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
)


Expand Down
97 changes: 71 additions & 26 deletions py/torch_tensorrt/dynamo/conversion/impl/arange.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -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")
34 changes: 33 additions & 1 deletion tests/py/dynamo/conversion/test_arange_aten.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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)
Expand Down
Loading