From 3883062f8a2d710ca1f5c179e5a340344929a1c3 Mon Sep 17 00:00:00 2001 From: gs-olive <113141689+gs-olive@users.noreply.github.com> Date: Thu, 7 Dec 2023 16:16:32 -0800 Subject: [PATCH] fix: Add support for Dynamic Shapes --- py/torch_tensorrt/dynamo/_DryRunTracker.py | 15 +++++++++++++++ py/torch_tensorrt/dynamo/_compiler.py | 12 ++++++++---- 2 files changed, 23 insertions(+), 4 deletions(-) diff --git a/py/torch_tensorrt/dynamo/_DryRunTracker.py b/py/torch_tensorrt/dynamo/_DryRunTracker.py index 031fce2e73..46d99ffe31 100644 --- a/py/torch_tensorrt/dynamo/_DryRunTracker.py +++ b/py/torch_tensorrt/dynamo/_DryRunTracker.py @@ -229,6 +229,21 @@ def input_formatter_helper(shapes: Any, dtypes: Any) -> str: if isinstance(shapes, tuple) and all(isinstance(elt, int) for elt in shapes): return f"Tensor: {shapes}@{str(dtypes)[6:]}, " + # Base case - dynamic shape, single dtype + elif ( + isinstance(shapes, dict) + and len(shapes) == 3 + and all( + ( + isinstance(shape, tuple) + and all(isinstance(elt, int) for elt in shape) + and k in ("min_shape", "opt_shape", "max_shape") + ) + for k, shape in shapes.items() + ) + ): + return f"Tensor: {shapes}@{str(dtypes)[6:]}, " + # Shapes is a sequence elif isinstance(shapes, (list, tuple)): formatted_str = "List[" if isinstance(shapes, list) else "Tuple(" diff --git a/py/torch_tensorrt/dynamo/_compiler.py b/py/torch_tensorrt/dynamo/_compiler.py index 23e32e2b65..ac7a323545 100644 --- a/py/torch_tensorrt/dynamo/_compiler.py +++ b/py/torch_tensorrt/dynamo/_compiler.py @@ -260,7 +260,7 @@ def compile_module( dryrun_tracker.total_ops_in_graph = total_ops dryrun_tracker.supported_ops_in_graph = num_supported_ops dryrun_tracker.graph_input_shapes = parse_complex_tensor_structs( - sample_inputs, "shape", tuple + sample_inputs, "shape", lambda x: dict(x) if isinstance(x, dict) else tuple(x) ) dryrun_tracker.graph_input_dtypes = parse_complex_tensor_structs( sample_inputs, "torch_dtype" @@ -372,7 +372,9 @@ def compile_module( ) subgraph_data.subgraph_input_shapes = parse_complex_tensor_structs( - submodule_inputs, "shape", tuple + submodule_inputs, + "shape", + lambda x: dict(x) if isinstance(x, dict) else tuple(x), ) subgraph_data.subgraph_input_dtypes = parse_complex_tensor_structs( submodule_inputs, "torch_dtype" @@ -383,7 +385,9 @@ def compile_module( ) subgraph_data.subgraph_output_shapes = parse_complex_tensor_structs( - submodule_outputs, "shape", tuple + submodule_outputs, + "shape", + lambda x: dict(x) if isinstance(x, dict) else tuple(x), ) subgraph_data.subgraph_output_dtypes = parse_complex_tensor_structs( submodule_outputs, "dtype" @@ -411,7 +415,7 @@ def compile_module( sample_outputs = [sample_outputs] dryrun_tracker.graph_output_shapes = parse_complex_tensor_structs( - sample_outputs, "shape", tuple + sample_outputs, "shape", lambda x: dict(x) if isinstance(x, dict) else tuple(x) ) dryrun_tracker.graph_output_dtypes = parse_complex_tensor_structs( sample_outputs, "dtype"