From cc17e5c5c67466351ab24609c6e98e549859b0f7 Mon Sep 17 00:00:00 2001 From: Joseph Loftin Date: Tue, 18 Aug 2026 21:43:43 +0000 Subject: [PATCH] Observable dry run --- py/torch_tensorrt/dynamo/_DryRunTracker.py | 2 + .../runtime/test_dryrun_stats_observable.py | 48 +++++++++++++++++++ 2 files changed, 50 insertions(+) create mode 100644 tests/py/dynamo/runtime/test_dryrun_stats_observable.py diff --git a/py/torch_tensorrt/dynamo/_DryRunTracker.py b/py/torch_tensorrt/dynamo/_DryRunTracker.py index 43789c4a0f..2007287b03 100644 --- a/py/torch_tensorrt/dynamo/_DryRunTracker.py +++ b/py/torch_tensorrt/dynamo/_DryRunTracker.py @@ -9,6 +9,7 @@ from torch_tensorrt.dynamo._settings import CompilationSettings from torch_tensorrt.dynamo.conversion._ConverterRegistry import ConverterRegistry from torch_tensorrt.dynamo.conversion.converter_utils import get_node_name +from torch_tensorrt.dynamo.observer import observable logger = logging.getLogger(__name__) @@ -67,6 +68,7 @@ class DryRunTracker: to_run_in_torch: List[str] = field(default_factory=list) +@observable() def dryrun_stats_display( dryrun_tracker: DryRunTracker, dryrun_enabled: Union[bool, str] ) -> None: diff --git a/tests/py/dynamo/runtime/test_dryrun_stats_observable.py b/tests/py/dynamo/runtime/test_dryrun_stats_observable.py new file mode 100644 index 0000000000..39d485f788 --- /dev/null +++ b/tests/py/dynamo/runtime/test_dryrun_stats_observable.py @@ -0,0 +1,48 @@ +# type: ignore + +import unittest + +import torch +import torch_tensorrt +from torch.testing._internal.common_utils import TestCase, run_tests +from torch_tensorrt.dynamo._DryRunTracker import DryRunTracker, dryrun_stats_display +from torch_tensorrt.dynamo.observer import ObserveContext + + +@unittest.skipIf(not torch.cuda.is_available(), "CUDA required") +class TestDryRunStatsObservable(TestCase): + def test_observer_captures_tracker(self): + class Add(torch.nn.Module): + def forward(self, x): + return x + x + + model = Add().cuda().eval() + x = torch.randn(2, 3, device="cuda") + trackers = [] + + def capture(ctx: ObserveContext) -> None: + trackers.append(ctx.args[0]) + + with dryrun_stats_display.observers.pre.add(capture): + torch_tensorrt.dynamo.compile( + model, + [x], + dryrun=True, + min_block_size=1, + enabled_precisions={torch.float32}, + ) + + self.assertEqual(len(trackers), 1) + tracker = trackers[0] + self.assertIsInstance(tracker, DryRunTracker) + self.assertGreater(tracker.total_ops_in_graph, 0) + self.assertGreaterEqual(tracker.supported_ops_in_graph, 0) + self.assertLessEqual( + tracker.supported_ops_in_graph, tracker.total_ops_in_graph + ) + self.assertGreaterEqual(tracker.tensorrt_graph_count, 1) + self.assertEqual(len(tracker.per_subgraph_data), tracker.tensorrt_graph_count) + + +if __name__ == "__main__": + run_tests()