diff --git a/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py b/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py index f287a1d1cf7..3250849d694 100644 --- a/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py +++ b/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py @@ -307,10 +307,14 @@ def partition_graph(self) -> torch.fx.GraphModule: Returns a GraphModule with submodules for each segment """ unsupported = getattr(self.operator_support, "unsupported_operators", None) + # unsupported_operators leaves out operators with side effects, so also check + # fallback_operators, which records every refusal including random and in-place ops. + # Without it a model that must run such an op in PyTorch would pass as fully supported. + fallback = getattr(self.operator_support, "fallback_operators", None) # The explicit assumption comes from the compiler's earlier support walk. # Otherwise, an empty dict means AccNodesFinder found no unsupported ops. fully_supported = self.assume_full_support or ( - isinstance(unsupported, dict) and len(unsupported) == 0 + isinstance(unsupported, dict) and len(unsupported) == 0 and not fallback ) # Fast path: user demanded a single TRT engine and every op is convertible. @@ -343,9 +347,13 @@ def partition_graph(self) -> torch.fx.GraphModule: # Delegate nodes based on operator coverage subgraphs = self.put_nodes_into_subgraphs() - # A graph is fully supported if there is a single partition and all operators are supported/convertible - full_support = len([s for s in subgraphs if s.is_acc]) == 1 and not getattr( - self.operator_support, "unsupported_operators", True + # A graph is fully supported if there is a single partition and all operators are + # supported/convertible. As above, unsupported_operators excludes side-effecting ops, + # so also require fallback_operators to be empty. + full_support = ( + len([s for s in subgraphs if s.is_acc]) == 1 + and not getattr(self.operator_support, "unsupported_operators", True) + and not getattr(self.operator_support, "fallback_operators", False) ) if not full_support and self.require_full_compilation: diff --git a/py/torch_tensorrt/dynamo/partitioning/_global_partitioner.py b/py/torch_tensorrt/dynamo/partitioning/_global_partitioner.py index 781ed9d1db5..2a2914269d3 100644 --- a/py/torch_tensorrt/dynamo/partitioning/_global_partitioner.py +++ b/py/torch_tensorrt/dynamo/partitioning/_global_partitioner.py @@ -69,9 +69,15 @@ def propose_partitions(self) -> List[Partition]: initial_proposed_partitions = super().propose_partitions() partitions = dict(enumerate(initial_proposed_partitions)) - # A graph is fully supported if there is a single partition and all operators are supported/convertible - full_support = len(partitions) == 1 and not getattr( - self.operator_support, "unsupported_operators", True + # A graph is fully supported if there is a single partition and all operators are + # supported/convertible. unsupported_operators does not include operators with side + # effects, such as random or in-place ops, so also check fallback_operators, which + # records every refusal including those. Otherwise a model that must run a random op + # in PyTorch would pass require_full_compilation. + full_support = ( + len(partitions) == 1 + and not getattr(self.operator_support, "unsupported_operators", True) + and not getattr(self.operator_support, "fallback_operators", False) ) if not full_support and self.require_full_compilation: diff --git a/py/torch_tensorrt/dynamo/partitioning/_hierarchical_partitioner.py b/py/torch_tensorrt/dynamo/partitioning/_hierarchical_partitioner.py index 85861adabd5..19aec81e31b 100644 --- a/py/torch_tensorrt/dynamo/partitioning/_hierarchical_partitioner.py +++ b/py/torch_tensorrt/dynamo/partitioning/_hierarchical_partitioner.py @@ -60,6 +60,10 @@ def __init__( # Initialize sets of supported/unsupported operators self.supported_operators: Dict[str, int] = {} self.unsupported_operators: Dict[str, int] = {} + # unsupported_operators skips operators with side effects, so it cannot tell whether + # a random or in-place op was refused. Record those here so require_full_compilation + # can reject them. + self.fallback_operators: Dict[str, int] = {} self.torch_executed_ops = torch_executed_ops # Map of backend names to sets of supported operators self.backend_support_map = backend_support_map @@ -96,6 +100,15 @@ def is_node_supported( self.unsupported_operators[node_name] = 1 else: self.unsupported_operators[node_name] += 1 + # Record impure refusals separately, since the gate above skips them. + if ( + i == len(self.backend_priority) - 1 + and node.is_impure() + and node.op in CALLABLE_NODE_OPS + ): + self.fallback_operators[node_name] = ( + self.fallback_operators.get(node_name, 0) + 1 + ) return False, NON_ACC_BACKEND_NAME @@ -248,9 +261,13 @@ def partition_graph(self) -> torch.fx.GraphModule: # Delegate nodes based on operator coverage subgraphs = self.put_nodes_into_subgraphs() - # A graph is fully supported if there is a single partition and all operators are supported/convertible - full_support = len([s for s in subgraphs if s.is_acc]) == 1 and not getattr( - self.operator_support, "unsupported_operators", True + # A graph is fully supported if there is a single partition and all operators are + # supported/convertible. unsupported_operators leaves out operators with side effects, + # so also check fallback_operators, which records refused random and in-place ops. + full_support = ( + len([s for s in subgraphs if s.is_acc]) == 1 + and not getattr(self.operator_support, "unsupported_operators", True) + and not getattr(self.operator_support, "fallback_operators", False) ) if not full_support and self.require_full_compilation: diff --git a/tests/py/dynamo/partitioning/test_000_full_support_detection.py b/tests/py/dynamo/partitioning/test_000_full_support_detection.py new file mode 100644 index 00000000000..63dcd362b34 --- /dev/null +++ b/tests/py/dynamo/partitioning/test_000_full_support_detection.py @@ -0,0 +1,165 @@ +import copy + +import torch +from parameterized import parameterized +from torch.testing._internal.common_utils import TestCase, run_tests +from torch_tensorrt.dynamo import partitioning +from torch_tensorrt.dynamo.lowering import ( + get_decompositions, + post_lowering, + pre_export_lowering, +) + +PARTITIONERS = [ + ("fast", partitioning.fast_partition), + ("global", partitioning.global_partition), + ("hierarchical", partitioning.hierarchical_adjacency_partition), +] + + +class TestFullSupportDetection(TestCase): + """A refused operator must never read as fully supported. + + An operator with side effects, such as a random or in-place op, that has no converter + is kept in PyTorch. unsupported_operators does not record it, since that dictionary + excludes impure nodes on purpose. fallback_operators records it instead, so + require_full_compilation must consult both when it decides whether a model is fully + supported. Otherwise a model that must run a random op in PyTorch would compile under + require_full_compilation=True. + """ + + @staticmethod + def _lower(module, args): + exported = torch.export.export(module.eval().cuda(), args) + lowered = exported.run_decompositions(get_decompositions(False)) + return post_lowering(pre_export_lowering(lowered).module()) + + @staticmethod + def _six_linear_layers(): + return torch.nn.ModuleList([torch.nn.Linear(64, 64) for _ in range(6)]) + + @classmethod + def _impure_refusal_module(cls): + class WithImpureRefusal(torch.nn.Module): + def __init__(self): + super().__init__() + self.layers = cls._six_linear_layers() + + def forward(self, x): + out = x + for index, layer in enumerate(self.layers): + out = torch.relu(layer(out)) + if index == 2: + # No converter, and impure, so it is kept in PyTorch. + out = out + torch.normal( + 0.0, + 1.0, + size=out.shape, + device=out.device, + dtype=out.dtype, + ) + return out + + return WithImpureRefusal() + + @classmethod + def _fully_supported_module(cls): + class FullySupported(torch.nn.Module): + def __init__(self): + super().__init__() + self.layers = cls._six_linear_layers() + + def forward(self, x): + out = x + for layer in self.layers: + out = torch.relu(layer(out)) + return out + + return FullySupported() + + @staticmethod + def _partition(partition_fn, graph_module, **kwargs): + if partition_fn is partitioning.hierarchical_adjacency_partition: + kwargs["backend_priority"] = ["tensorrt"] + # Both partitioners mutate the module they are given, so hand each a copy. + return partition_fn(copy.deepcopy(graph_module), **kwargs) + + @parameterized.expand(PARTITIONERS) + def test_impure_refusal_is_not_fully_supported(self, _, partition_fn): + graph_module = self._lower( + self._impure_refusal_module(), (torch.randn(8, 64, device="cuda"),) + ) + with self.assertRaisesRegex(AssertionError, "not fully supported"): + self._partition( + partition_fn, + graph_module, + min_block_size=1, + require_full_compilation=True, + ) + + @parameterized.expand(PARTITIONERS) + def test_impure_refusal_is_recorded_as_fallback(self, _, partition_fn): + """The refused impure operator is recorded in fallback_operators, not in + unsupported_operators. unsupported_operators keeps excluding impure nodes, which is + the contract the fallback reporting relies on.""" + graph_module = self._lower( + self._impure_refusal_module(), (torch.randn(8, 64, device="cuda"),) + ) + _, support = self._partition(partition_fn, graph_module, min_block_size=1) + self.assertTrue( + support.fallback_operators, + "the refused impure operator was not recorded in fallback_operators, so " + "require_full_compilation cannot see it", + ) + + @parameterized.expand(PARTITIONERS) + def test_pure_refusal_is_not_fully_supported(self, _, partition_fn): + """A refused pure operator was already caught. Keep it that way.""" + graph_module = self._lower( + self._fully_supported_module(), (torch.randn(8, 64, device="cuda"),) + ) + with self.assertRaisesRegex(AssertionError, "not fully supported"): + self._partition( + partition_fn, + graph_module, + min_block_size=1, + require_full_compilation=True, + torch_executed_ops={"torch.ops.aten.relu.default"}, + ) + + @parameterized.expand(PARTITIONERS) + def test_fully_supported_module_is_accepted(self, name, partition_fn): + """Guards against over correction. This passes before the change too, so it does + not prove the fix; it proves the fix did not start rejecting good graphs.""" + graph_module = self._lower( + self._fully_supported_module(), (torch.randn(8, 64, device="cuda"),) + ) + partitioned, support = self._partition( + partition_fn, + graph_module, + min_block_size=1, + require_full_compilation=True, + ) + self.assertFalse(support.unsupported_operators) + self.assertFalse(support.fallback_operators) + blocks = [child for child, _ in partitioned.named_children()] + self.assertTrue( + any("_run_on_acc" in block for block in blocks), + f"expected an accelerated block, got {blocks}", + ) + # The global partitioner leaves a refused node inline in the parent rather than + # naming a torch block, so assert on what is left outside the block instead. + remaining = [ + node.name + for node in partitioned.graph.nodes + if node.op == "call_function" and "_run_on_acc" not in node.name + ] + self.assertEqual( + remaining, + [], + f"expected no operator left outside the engine, got {remaining}", + ) + + +if __name__ == "__main__": + run_tests()