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
71 changes: 32 additions & 39 deletions py/torch_tensorrt/dynamo/conversion/_ConverterRegistry.py
Original file line number Diff line number Diff line change
Expand Up @@ -138,63 +138,56 @@ def _has_dynamic_shapes(
node, torch.fx.Node
), "Inputs to validator functions must be FX Nodes"

def _contains_dynamic_value(value: Any) -> bool:
"""Recursively checks FX metadata for symbolic values or dimensions."""
if isinstance(value, (SymFloat, SymInt, SymBool)):
return True
if isinstance(value, dict):
return any(_contains_dynamic_value(item) for item in value.values())
if isinstance(value, (list, tuple)):
return any(_contains_dynamic_value(item) for item in value)
if value is None or isinstance(value, (int, float, bool)):
return False
if hasattr(value, "size"):
return any(_contains_dynamic_value(dim) for dim in value.size())
return False

def _is_subnode_dynamic(subnode: torch.fx.Node) -> bool:
"""Checks if a node itself has Dynamic properties"""
_has_symbolic_sizes_strides, is_shape_dynamic = False, False
if "val" in subnode.meta:
_has_symbolic_sizes_strides = getattr(
subnode.meta["val"], "_has_symbolic_sizes_strides", False
)
meta_val = subnode.meta["val"]
if isinstance(meta_val, (list, tuple)):
for val in meta_val:
if val is None:
continue
if isinstance(val, (SymFloat, SymInt, SymBool)):
is_shape_dynamic = True
break
if not hasattr(val, "size"):
continue
shape = val.size()
if any(
isinstance(dim, (SymFloat, SymInt, SymBool)) for dim in shape
):
is_shape_dynamic = True
break
elif isinstance(meta_val, (SymFloat, SymInt, SymBool)):
is_shape_dynamic = True
elif isinstance(meta_val, (int, float, bool)):
is_shape_dynamic = False
else:
shape = subnode.meta["val"].size()
is_shape_dynamic = any(
isinstance(dim, (SymFloat, SymInt, SymBool)) for dim in shape
)
"""Checks if a node itself has Dynamic properties."""
meta_val = subnode.meta.get("val", subnode.meta.get("example_value"))
return bool(
getattr(meta_val, "_has_symbolic_sizes_strides", False)
or _contains_dynamic_value(meta_val)
)

return _has_symbolic_sizes_strides or is_shape_dynamic
def _is_argument_dynamic(argument: Argument) -> bool:
"""Checks nodes nested inside FX container arguments."""
if isinstance(argument, torch.fx.Node):
return _is_subnode_dynamic(argument)
if isinstance(argument, dict):
return any(_is_argument_dynamic(item) for item in argument.values())
if isinstance(argument, (list, tuple)):
return any(_is_argument_dynamic(item) for item in argument)
return False

# Check node value itself
if arg_positions_to_check is None and _is_subnode_dynamic(node):
return True

# Check node arguments individually
if arg_positions_to_check is None and any(
_is_subnode_dynamic(arg) for arg in node.args if isinstance(arg, torch.fx.Node)
_is_argument_dynamic(arg) for arg in node.args
):
return True
# Check specific arg positions if the caller has specified positions to check
elif arg_positions_to_check is not None and any(
_is_subnode_dynamic(node.args[i])
for i in arg_positions_to_check
if isinstance(node.args[i], torch.fx.Node)
_is_argument_dynamic(node.args[i]) for i in arg_positions_to_check
):
return True

# Check node keyword arguments individually
if arg_positions_to_check is None and any(
_is_subnode_dynamic(kwarg)
for kwarg in node.kwargs.values()
if isinstance(kwarg, torch.fx.Node)
_is_argument_dynamic(kwarg) for kwarg in node.kwargs.values()
):
return True

Expand Down
10 changes: 9 additions & 1 deletion py/torch_tensorrt/dynamo/conversion/impl/arange.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ def _sequence_dtype(
if isinstance(x, float):
return trt.DataType.FLOAT

return trt.DataType.INT64
return trt.DataType.INT64


def arange(
Expand All @@ -65,6 +65,9 @@ def arange(
start_rank_0 = get_trt_tensor(
ctx, start, name + "_start_rank_0", value_dtype, min_rank=0
)
start_rank_0 = cast_trt_tensor(
ctx, start_rank_0, value_dtype, name + "_start_rank_0_casted"
)
# 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 len(start_rank_0.shape) > 0:
Expand All @@ -80,6 +83,11 @@ def arange(
)
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)
start_rank_1 = cast_trt_tensor(
ctx, start_rank_1, value_dtype, name + "_start_rank_1_casted"
)
end = cast_trt_tensor(ctx, end, value_dtype, name + "_end_casted")
step = cast_trt_tensor(ctx, step, value_dtype, name + "_step_casted")

# The number of elements is ceil((end - start) / step), computed as
# -floor((start - end) / step) so that the whole expression stays in the
Expand Down
35 changes: 33 additions & 2 deletions tests/py/dynamo/conversion/test_arange_aten.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
import unittest

import torch
import torch.nn as nn
import torch_tensorrt
Expand Down Expand Up @@ -93,6 +91,39 @@ def forward(self, end_tensor):
use_dynamo_tracer=False,
)

def test_arange_data_dependent_start(self):
class Arange(torch.nn.Module):
def forward(self, x, mask):
end = mask.nonzero().size(0)
start = end // 2
indices = torch.arange(start, end, 1, device=x.device)
return x.index_select(0, indices)

previous_capture_setting = torch._dynamo.config.capture_dynamic_output_shape_ops
try:
torch._dynamo.config.capture_dynamic_output_shape_ops = True
torch._dynamo.reset()

model = Arange().eval().cuda()
x = torch.randn((16, 8), device="cuda")
mask = torch.arange(16, device="cuda") % 2 == 0
expected = model(x, mask)

compiled = torch.compile(
model,
backend="tensorrt",
options={
"pass_through_build_failures": True,
"min_block_size": 1,
},
)
torch.testing.assert_close(compiled(x, mask), expected)
finally:
torch._dynamo.config.capture_dynamic_output_shape_ops = (
previous_capture_setting
)
torch._dynamo.reset()


if __name__ == "__main__":
run_tests()
31 changes: 31 additions & 0 deletions tests/py/dynamo/conversion/test_split_aten.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@
from torch.testing._internal.common_utils import run_tests
from torch_tensorrt import Input
from torch_tensorrt.dynamo.conversion import UnsupportedOperatorException
from torch_tensorrt.dynamo.conversion._ConverterRegistry import (
has_static_shapes_in_args,
)

from .harness import DispatchTestCase

Expand Down Expand Up @@ -192,6 +195,34 @@ def forward(self, input):
input,
)

def test_dynamic_split_sections_decline_conversion(self):
columns = 8

class RuntimeSections(torch.nn.Module):
def forward(self, x, k):
first_size = k.item()
torch._check(first_size >= 1)
torch._check(first_size <= columns - 1)
return torch.split(x, [first_size, columns - first_size], dim=1)

class LiteralSections(torch.nn.Module):
def forward(self, x, k):
return torch.split(x, [3, columns - 3], dim=1)

def get_split_node(model, args):
exported = torch.export.export(model, args)
return next(
node
for node in exported.graph.nodes
if node.target == torch.ops.aten.split_with_sizes.default
)

args = (torch.rand(4, columns), torch.tensor([3]))
validator = has_static_shapes_in_args([1])

self.assertFalse(validator(get_split_node(RuntimeSections(), args), None))
self.assertTrue(validator(get_split_node(LiteralSections(), args), None))


if __name__ == "__main__":
run_tests()
Loading