Skip to content
Merged
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
16 changes: 12 additions & 4 deletions py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down
12 changes: 9 additions & 3 deletions py/torch_tensorrt/dynamo/partitioning/_global_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
23 changes: 20 additions & 3 deletions py/torch_tensorrt/dynamo/partitioning/_hierarchical_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down
165 changes: 165 additions & 0 deletions tests/py/dynamo/partitioning/test_000_full_support_detection.py
Original file line number Diff line number Diff line change
@@ -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()
Loading