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
42 changes: 41 additions & 1 deletion py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py
Original file line number Diff line number Diff line change
Expand Up @@ -1404,7 +1404,47 @@ def aten_ops_clamp(
)


@dynamo_tensorrt_converter(torch.ops.aten.gather.default)
def gather_validator(
node: Node, settings: Optional[CompilationSettings] = None
) -> bool:
"""Keep cases the TensorRT gather cannot serve on the PyTorch path.

An empty index gives the engine a zero size output binding, which TensorRT refuses to
enqueue. The call still returns, so every other output of that engine silently comes
back zeroed. An engine built for a float64 input expects float32 and rejects the
caller's tensor at run time unless truncate_double is set, and a uint8 output fails the
engine build outright. All three ran correctly in PyTorch before this converter
accepted dynamic shapes.
"""
data_meta = (
node.args[0].meta.get("tensor_meta") if hasattr(node.args[0], "meta") else None
)
index_meta = (
node.args[2].meta.get("tensor_meta") if hasattr(node.args[2], "meta") else None
)
if index_meta is not None and 0 in tuple(index_meta.shape):
_LOGGER.debug("gather with an empty index is not supported, falling back")
return False
if data_meta is None:
return True
if data_meta.dtype == torch.uint8:
_LOGGER.debug("gather with a uint8 input is not supported, falling back")
return False
if data_meta.dtype == torch.float64 and not (
settings is not None and settings.truncate_double
):
_LOGGER.debug(
"gather with a float64 input needs truncate_double=True, falling back"
)
return False
return True


@dynamo_tensorrt_converter(
torch.ops.aten.gather.default,
capability_validator=gather_validator,
supports_dynamic_shapes=True,
)
@enforce_tensor_types(
{
0: (TRTTensor,),
Expand Down
67 changes: 67 additions & 0 deletions tests/py/dynamo/conversion/test_gather_aten.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,3 +73,70 @@ def forward(self, input, index):
input = torch.zeros(3, 5, dtype=torch.int32)
inputs = [input, index]
self.run_test(TestModule(), inputs)

@parameterized.expand(
[
("positive_dim", 1),
("negative_dim", -1),
]
)
def test_gather_dynamic_shape(self, _, dim):
"""The registry only consults supports_dynamic_shapes when a node carries symbolic
shape metadata, which the legacy tracer does not produce, so this needs the dynamo
tracer or it passes either way."""

class TestModule(torch.nn.Module):
def forward(self, input, index):
return torch.ops.aten.gather.default(input, dim, index)

input_specs = [
Input(
min_shape=(1, 5),
opt_shape=(3, 5),
max_shape=(6, 5),
dtype=torch.float32,
),
Input(
min_shape=(1, 4),
opt_shape=(3, 4),
max_shape=(6, 4),
dtype=torch.int64,
),
]
self.run_test_with_dynamic_shape(
TestModule(),
input_specs,
use_dynamo_tracer=True,
)

def test_gather_dynamic_gathered_axis(self):
"""The gathered axis itself is dynamic here, so index validity depends on the shape
the engine is given at run time rather than on the shape it was built at."""

class TestModule(torch.nn.Module):
def forward(self, input, index):
return torch.ops.aten.gather.default(input, 1, index)

input_specs = [
Input(
min_shape=(3, 1),
opt_shape=(3, 4),
max_shape=(3, 6),
dtype=torch.float32,
),
Input(
min_shape=(3, 1),
opt_shape=(3, 4),
max_shape=(3, 6),
dtype=torch.int64,
),
]
self.run_test_with_dynamic_shape(
TestModule(),
input_specs,
use_dynamo_tracer=True,
)


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