diff --git a/py/torch_tensorrt/dynamo/conversion/_conversion.py b/py/torch_tensorrt/dynamo/conversion/_conversion.py index 38d104473a..ccc29c76e3 100644 --- a/py/torch_tensorrt/dynamo/conversion/_conversion.py +++ b/py/torch_tensorrt/dynamo/conversion/_conversion.py @@ -221,7 +221,9 @@ def interpret_module_to_result( SerializedInterpreterResult """ - symbolic_shape_expressions = extract_symbolic_shape_expressions(module) + symbolic_shape_expressions = extract_symbolic_shape_expressions( + module, truncate_double=settings.truncate_double + ) if symbolic_shape_expressions is None: raise RuntimeError( "Failed to extract symbolic shape expressions from source FX graph partition" diff --git a/py/torch_tensorrt/dynamo/conversion/_symbolic_shape_capture.py b/py/torch_tensorrt/dynamo/conversion/_symbolic_shape_capture.py index e8f84921a2..a952703435 100644 --- a/py/torch_tensorrt/dynamo/conversion/_symbolic_shape_capture.py +++ b/py/torch_tensorrt/dynamo/conversion/_symbolic_shape_capture.py @@ -16,6 +16,7 @@ def extract_symbolic_shape_expressions( module: torch.fx.GraphModule, + truncate_double: bool = False, ) -> Optional[Dict[str, List[Dict[str, Any]]]]: """ Extract symbolic shape expressions from an FX graph. @@ -25,6 +26,8 @@ def extract_symbolic_shape_expressions( Args: module: FX GraphModule with symbolic shapes in node metadata + truncate_double: Record float64 tensor bindings as float32, matching + the precision TensorRT builds when double truncation is enabled Returns: Dict with 'inputs' and 'outputs' keys, each containing a list of dicts with shape_exprs and dtype, @@ -64,7 +67,11 @@ def extract_symbolic_shape_expressions( input_info.append( { "shape_exprs": shape_exprs, - "dtype": input_val.dtype, + "dtype": ( + torch.float32 + if truncate_double and input_val.dtype == torch.float64 + else input_val.dtype + ), "name": input_node.name, } ) @@ -115,7 +122,11 @@ def extract_symbolic_shape_expressions( output_info.append( { "shape_exprs": shape_exprs, - "dtype": out_val.dtype, + "dtype": ( + torch.float32 + if truncate_double and out_val.dtype == torch.float64 + else out_val.dtype + ), } ) elif isinstance(out_val, (torch.SymInt, torch.SymFloat, int, float, bool)): diff --git a/py/torch_tensorrt/dynamo/conversion/truncate_double.py b/py/torch_tensorrt/dynamo/conversion/truncate_double.py index 51e35a7840..80d5aea93d 100644 --- a/py/torch_tensorrt/dynamo/conversion/truncate_double.py +++ b/py/torch_tensorrt/dynamo/conversion/truncate_double.py @@ -1,13 +1,13 @@ from __future__ import annotations import logging -from typing import Optional, Sequence, Set +from typing import Any, Dict, Optional, Sequence, Set import torch from torch.fx.node import _get_qualified_name from torch_tensorrt._enums import dtype from torch_tensorrt._Input import Input -from torch_tensorrt.dynamo.utils import get_torch_inputs +from torch_tensorrt.dynamo.utils import get_output_metadata, get_torch_inputs logger = logging.getLogger(__name__) @@ -40,46 +40,58 @@ def _extract_downstream_get_nodes( return get_nodes +def _metadata_dtype(metadata: Dict[str, Any]) -> Optional[torch.dtype]: + """Return the dtype of tensor metadata, ignoring scalar outputs.""" + value = metadata.get("val") + if isinstance(value, torch.Tensor): + return value.dtype + + tensor_meta = metadata.get("tensor_meta") + return getattr(tensor_meta, "dtype", None) + + +def _metadata_to_dtype( + metadata: Dict[str, Any], target_dtype: torch.dtype +) -> Dict[str, Any]: + """Copy tensor metadata while changing its dtype.""" + updated = metadata.copy() + value = updated.get("val") + if isinstance(value, torch.Tensor): + updated["val"] = value.to(target_dtype) + + tensor_meta = updated.get("tensor_meta") + if tensor_meta is not None and hasattr(tensor_meta, "_replace"): + updated["tensor_meta"] = tensor_meta._replace(dtype=target_dtype) + + return updated + + def _repair_64bit_input( gm: torch.fx.GraphModule, position: int, submodule_name: str, - submodule_outputs: Optional[torch.Tensor | Sequence[torch.Tensor]], + submodule_output_metadata: Optional[Sequence[Dict[str, Any]]], + is_collection_output: bool, dtype: torch.dtype, ) -> None: - """Fixes a single Long/Double input to a TRT-accelerated subgraph - - In-Place modifies the provided graph - - Inserts a cast to the 32-bit equivalent type for TRT, then if necessary, - inserts an upcast back to the 64-bit type for subsequent Torch operations + """Fix a single double input and any double outputs at a TRT boundary. - Args: - gm: FX GraphModule enclosing the TRT subgraph - position: Index in the submodule inputs at which the long or double input is found - submodule_name: Name of TRT-accelerated subgraph module in FX graph - submodule_outputs: Output tensor(s) of TRT-accelerated subgraph (used for dtypes/structure) - dtype: Data type of tensor at position in submodule (double/long) + The output dtypes come from the partition's FX metadata. Compilation must not + execute the partition merely to discover information already recorded there. """ - assert dtype in ( - torch.float64, - ), f"dtype argument must be torch.float64, got {dtype}" + assert dtype == torch.float64, f"dtype argument must be torch.float64, got {dtype}" logger.info( f"Downcasting a 64-bit input at position {position} of submodule {submodule_name}" ) - # Determine target data type in 32 and 64 bit forms dtype_64bit = dtype dtype_32bit = torch.float32 - # Find the node representing the submodule in the graph module_node = None - - # Iterate over all nodes in the graph, seeking target module name match - for n in gm.graph.nodes: - if n.op == "call_module" and str(n.target) == submodule_name: - module_node = n + for node in gm.graph.nodes: + if node.op == "call_module" and str(node.target) == submodule_name: + module_node = node break if module_node is None: @@ -87,77 +99,81 @@ def _repair_64bit_input( f"Sought module node {submodule_name}, could not find in graph:\n{gm.graph}" ) - # Extract the 64-bit node of the input node_64bit = module_node.all_input_nodes[position] - - # Prior to the module, insert a cast to the 32-bit equivalent node with gm.graph.inserting_before(module_node): node_32bit = gm.graph.call_function( torch.ops.aten._to_copy.default, args=(node_64bit,), kwargs={"dtype": dtype_32bit}, ) + node_32bit.meta = _metadata_to_dtype(node_64bit.meta, dtype_32bit) - # Replace 64-bit input to TRT module with new 32-bit cast node module_node.replace_input_with(node_64bit, node_32bit) - output_positions_64bit = set() - - # Determine if any outputs of the model are 64-bit type and store their indices - if submodule_outputs is not None: - outputs_list = ( - [submodule_outputs] - if isinstance(submodule_outputs, torch.Tensor) - else submodule_outputs - ) - - for output_position, output in enumerate(outputs_list): - if output.dtype == dtype_64bit: - output_positions_64bit.add(output_position) + output_positions_64bit: Set[int] = set() + original_output_metadata = list(submodule_output_metadata or ()) + truncated_output_metadata = [] + for output_position, metadata in enumerate(original_output_metadata): + if _metadata_dtype(metadata) == dtype_64bit: + output_positions_64bit.add(output_position) + truncated_output_metadata.append(_metadata_to_dtype(metadata, dtype_32bit)) + else: + truncated_output_metadata.append(metadata.copy()) + + # The call_module node describes the actual engine boundary. Preserve its + # container convention while correcting tensor dtypes to what TRT emits. + if truncated_output_metadata: + for key in ("val", "tensor_meta"): + values = [ + metadata[key] + for metadata in truncated_output_metadata + if key in metadata + ] + if not values: + continue + current = module_node.meta.get(key) + if isinstance(current, tuple): + module_node.meta[key] = tuple(values) + elif isinstance(current, list) or len(values) > 1: + module_node.meta[key] = values + else: + module_node.meta[key] = values[0] - # Only enter this code block if there exists a 64-bit output - # This implies a cast is needed, since TRT cannot output 64-bit tensors if output_positions_64bit: - # Determine whether the outputs of the module are tuple-type or not - is_collection_output = False - if isinstance(submodule_outputs, tuple): - is_collection_output = True - if not is_collection_output: - # If the output is a single tensor, insert a cast back to int64 with gm.graph.inserting_after(module_node): cast_node_64bit = gm.graph.call_function( torch.ops.aten._to_copy.default, args=(module_node,), kwargs={"dtype": dtype_64bit}, ) + cast_node_64bit.meta = original_output_metadata[0].copy() - # Replace all uses of the TRT module (except the cast node) with the 64-bit equivalent module_node.replace_all_uses_with( - cast_node_64bit, delete_user_cb=lambda user: (user != cast_node_64bit) + cast_node_64bit, delete_user_cb=lambda user: user != cast_node_64bit ) - else: - # If the output is a tuple of tensors, extract downstream users for each 64-bit output get_nodes = _extract_downstream_get_nodes( module_node, output_positions_64bit ) - - # For each downstream user, append a cast node back to the 64-bit precision for get_node in get_nodes: + output_position = get_node.args[1] + get_node.meta = truncated_output_metadata[output_position].copy() with gm.graph.inserting_after(get_node): cast_node_64bit = gm.graph.call_function( torch.ops.aten._to_copy.default, args=(get_node,), - kwargs={"dtype": torch.float64}, + kwargs={"dtype": dtype_64bit}, ) + cast_node_64bit.meta = original_output_metadata[ + output_position + ].copy() get_node.replace_all_uses_with( cast_node_64bit, - delete_user_cb=lambda user: (user != cast_node_64bit), + delete_user_cb=lambda user: user != cast_node_64bit, ) - # Clean up graph and ensure invariants are preserved gm.graph.eliminate_dead_code() gm.graph.lint() gm.recompile() @@ -188,24 +204,21 @@ def repair_double_inputs( submodule_torch_inputs = get_torch_inputs(submodule_inputs, device) num_submodule_inputs = len(submodule_inputs) repaired_outputs_once = False + output_node = next(node for node in submodule.graph.nodes if node.op == "output") + is_collection_output = isinstance(output_node.args[0], (tuple, list)) + submodule_output_metadata = get_output_metadata(submodule) - # For each input to the TRT subgraph, check if its type is long/double + # For each input to the TRT subgraph, check if its type is double. for position in range(num_submodule_inputs): param = submodule_torch_inputs[position] - # If the data type of the input is long/double, insert necessary - # casts to replace the operation if isinstance(param, torch.Tensor) and param.dtype == torch.float64: - # Ensure outputs are only repaired once per submodule to avoid - # unnecessary ops showing up in the graph - if not repaired_outputs_once: - submodule_outputs = submodule(*submodule_torch_inputs) - _repair_64bit_input( parent_graph, position, submodule_name if submodule_name is not None else submodule._get_name(), - None if repaired_outputs_once else submodule_outputs, + None if repaired_outputs_once else submodule_output_metadata, + is_collection_output, param.dtype, ) diff --git a/tests/py/dynamo/conversion/test_truncate_double.py b/tests/py/dynamo/conversion/test_truncate_double.py new file mode 100644 index 0000000000..d96a718d78 --- /dev/null +++ b/tests/py/dynamo/conversion/test_truncate_double.py @@ -0,0 +1,143 @@ +import operator +from unittest.mock import Mock + +import torch +from torch import nn +from torch.testing._internal.common_utils import TestCase, run_tests +from torch_tensorrt import Input +from torch_tensorrt._enums import dtype +from torch_tensorrt.dynamo.conversion._symbolic_shape_capture import ( + extract_symbolic_shape_expressions, +) +from torch_tensorrt.dynamo.conversion.truncate_double import repair_double_inputs + + +class TestTruncateDoubleMetadata(TestCase): + def _make_graphs( + self, *, with_scalar_output: bool + ) -> tuple[torch.fx.GraphModule, torch.fx.GraphModule, torch.fx.Node]: + tensor64 = torch.empty((2, 3), dtype=torch.float64) + + subgraph = torch.fx.Graph() + subgraph_input = subgraph.placeholder("x") + subgraph_input.meta["val"] = tensor64 + tensor_output = subgraph.call_function( + torch.ops.aten.add.Tensor, args=(subgraph_input, 1.0) + ) + tensor_output.meta["val"] = tensor64 + if with_scalar_output: + scalar_output = subgraph.call_function( + torch.ops.aten.sym_size.int, args=(tensor_output, 0) + ) + scalar_output.meta["val"] = 2 + subgraph.output((tensor_output, scalar_output)) + else: + subgraph.output(tensor_output) + submodule = torch.fx.GraphModule({}, subgraph) + + root = nn.Module() + root.add_module("run_on_acc_0", submodule) + parent_graph = torch.fx.Graph() + parent_input = parent_graph.placeholder("x") + parent_input.meta["val"] = tensor64 + engine_node = parent_graph.call_module("run_on_acc_0", args=(parent_input,)) + if with_scalar_output: + engine_node.meta["val"] = [tensor64, 2] + tensor_getitem = parent_graph.call_function( + operator.getitem, args=(engine_node, 0) + ) + tensor_getitem.meta["val"] = tensor64 + scalar_getitem = parent_graph.call_function( + operator.getitem, args=(engine_node, 1) + ) + scalar_getitem.meta["val"] = 2 + parent_graph.output((tensor_getitem, scalar_getitem)) + else: + engine_node.meta["val"] = tensor64 + parent_graph.output(engine_node) + + parent = torch.fx.GraphModule(root, parent_graph) + return parent, submodule, engine_node + + def _repair( + self, parent: torch.fx.GraphModule, submodule: torch.fx.GraphModule + ) -> Input: + # Any attempt to rediscover output dtypes by executing the partition is a + # regression: parameters may be offloaded and outputs may include scalars. + submodule.forward = Mock( + side_effect=AssertionError("truncate_double executed the partition") + ) + input_spec = Input((2, 3), dtype=torch.float64) + repaired = repair_double_inputs( + parent, + submodule, + [input_spec], + torch.device("cpu"), + "run_on_acc_0", + ) + return repaired[0] + + def test_single_output_metadata_matches_engine_boundary(self): + parent, submodule, engine_node = self._make_graphs(with_scalar_output=False) + + repaired_input = self._repair(parent, submodule) + + casts = [ + node + for node in parent.graph.nodes + if node.target == torch.ops.aten._to_copy.default + ] + self.assertEqual(len(casts), 2) + input_cast = next( + node for node in casts if node.kwargs["dtype"] == torch.float32 + ) + output_cast = next( + node for node in casts if node.kwargs["dtype"] == torch.float64 + ) + self.assertIs(engine_node.args[0], input_cast) + self.assertEqual(input_cast.meta["val"].dtype, torch.float32) + self.assertEqual(engine_node.meta["val"].dtype, torch.float32) + self.assertEqual(output_cast.meta["val"].dtype, torch.float64) + self.assertEqual(repaired_input.dtype, dtype.float32) + + def test_scalar_tuple_output_is_not_executed_or_retyped(self): + parent, submodule, engine_node = self._make_graphs(with_scalar_output=True) + + self._repair(parent, submodule) + + engine_values = engine_node.meta["val"] + self.assertEqual(engine_values[0].dtype, torch.float32) + self.assertEqual(engine_values[1], 2) + + getitems = [ + node for node in parent.graph.nodes if node.target == operator.getitem + ] + tensor_getitem = next(node for node in getitems if node.args[1] == 0) + scalar_getitem = next(node for node in getitems if node.args[1] == 1) + self.assertEqual(tensor_getitem.meta["val"].dtype, torch.float32) + self.assertEqual(scalar_getitem.meta["val"], 2) + + restoring_casts = [ + node + for node in parent.graph.nodes + if node.target == torch.ops.aten._to_copy.default + and node.kwargs["dtype"] == torch.float64 + ] + self.assertEqual(len(restoring_casts), 1) + self.assertIs(restoring_casts[0].args[0], tensor_getitem) + self.assertEqual(restoring_casts[0].meta["val"].dtype, torch.float64) + + def test_symbolic_shape_metadata_uses_truncated_binding_dtype(self): + _, submodule, _ = self._make_graphs(with_scalar_output=True) + + metadata = extract_symbolic_shape_expressions(submodule, truncate_double=True) + + self.assertIsNotNone(metadata) + self.assertEqual(metadata["inputs"][0]["dtype"], torch.float32) + self.assertEqual(metadata["outputs"][0]["dtype"], torch.float32) + self.assertEqual(metadata["outputs"][1]["dtype"], torch.int64) + self.assertTrue(metadata["outputs"][1]["is_scalar"]) + + +if __name__ == "__main__": + run_tests()