Skip to content
Open
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
2 changes: 2 additions & 0 deletions py/torch_tensorrt/dynamo/_DryRunTracker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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:
Expand Down
48 changes: 48 additions & 0 deletions tests/py/dynamo/runtime/test_dryrun_stats_observable.py
Original file line number Diff line number Diff line change
@@ -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(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

dynamo.compile() needs an ExportedProgram, not an nn.Module. Export first with torch.export.export(model, (x,)), then pass that into compile.

Also drop enabled_precisions={torch.float32}, enabled_precisions is deprecated as of TRT 11. Dynamo uses dtypes on the model and input so this FP32 test needs no precision argument.

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()
Loading