From 7762e49c850857fc271b04fd826e012770cf1e08 Mon Sep 17 00:00:00 2001 From: Michael Williams Date: Wed, 26 Aug 2026 15:40:45 -0700 Subject: [PATCH] Skip rebuilding FX forward after post-lowering and before adjacency split. TRTInterpreter and AccNodesFinder walk graph.nodes; split still rebuilds forward when it constructs the runtime GraphModules. --- .../lowering/passes/_aten_lowering_pass.py | 17 ++- .../passes/batch_cheap_fx_cleanups.py | 4 +- .../dynamo/lowering/passes/pass_utils.py | 24 +++- .../partitioning/_adjacency_partitioner.py | 5 +- tests/py/dynamo/conversion/harness.py | 4 +- .../test_skip_conversion_recompile.py | 104 ++++++++++++++++++ tools/llm/tests/test_static_cache.py | 2 +- 7 files changed, 145 insertions(+), 15 deletions(-) create mode 100644 tests/py/dynamo/lowering/test_skip_conversion_recompile.py diff --git a/py/torch_tensorrt/dynamo/lowering/passes/_aten_lowering_pass.py b/py/torch_tensorrt/dynamo/lowering/passes/_aten_lowering_pass.py index e829f6e7b61..409b38e6e21 100644 --- a/py/torch_tensorrt/dynamo/lowering/passes/_aten_lowering_pass.py +++ b/py/torch_tensorrt/dynamo/lowering/passes/_aten_lowering_pass.py @@ -160,20 +160,29 @@ def get_lowering_pass_config(lowering_pass: LoweringPassSignature) -> dict[str, def post_lowering( - gm: torch.fx.GraphModule, settings: CompilationSettings = CompilationSettings() + gm: torch.fx.GraphModule, + settings: CompilationSettings = CompilationSettings(), + *, + recompile: bool = False, ) -> torch.fx.GraphModule: - """Applies the lowering passes to a graph module after torch.export/torch.compile and their decompositions, returns the modified GraphModule""" + """Applies the lowering passes to a graph module after torch.export/torch.compile and their decompositions, returns the modified GraphModule + + By default the trailing cleanup skips ``gm.recompile()``. Conversion walks + ``gm.graph.nodes`` (TRTInterpreter) and does not call ``gm.forward()``. + Pass ``recompile=True`` when the returned module will be executed as Python + (tests, dry-run, Python fallback). + """ logging.debug( f"Invoking DynamoPassManager and applying lowering passes: {ATEN_POST_LOWERING_PASSES}" ) fake_mode = torch._export.utils._detect_fake_mode_from_gm(gm) fake_tensor_updater = FakeTensorUpdater(gm) - # Batch DCE/lint/recompile across passes: each pass may still call + # Batch DCE/lint across passes: each pass may still call # clean_up_graph_after_modifications, but only the final flush pays for it. set_defer_graph_cleanup(True) try: gm = ATEN_POST_LOWERING_PASSES(gm, settings) - gm = flush_deferred_graph_cleanup(gm) + gm = flush_deferred_graph_cleanup(gm, recompile=recompile) finally: set_defer_graph_cleanup(False) if fake_mode is not None: diff --git a/py/torch_tensorrt/dynamo/lowering/passes/batch_cheap_fx_cleanups.py b/py/torch_tensorrt/dynamo/lowering/passes/batch_cheap_fx_cleanups.py index c6794ed19c6..03a95d1879c 100644 --- a/py/torch_tensorrt/dynamo/lowering/passes/batch_cheap_fx_cleanups.py +++ b/py/torch_tensorrt/dynamo/lowering/passes/batch_cheap_fx_cleanups.py @@ -1,8 +1,8 @@ """Batch cheap, non-conflicting FX cleanups into one graph repair cycle. Several post-lowering passes only delete/rewrite a few node kinds, then each -calls ``clean_up_graph_after_modifications`` (DCE + lint + recompile). On large -graphs that recompile dominates and is paid repeatedly. +calls ``clean_up_graph_after_modifications`` (DCE + lint, and recompile when +not deferred). On large graphs that recompile dominates and is paid repeatedly. Naren's guidance: use one iteration / one cleanup for repairs that do not conflict. This pass runs those mutations back-to-back and cleans up once. diff --git a/py/torch_tensorrt/dynamo/lowering/passes/pass_utils.py b/py/torch_tensorrt/dynamo/lowering/passes/pass_utils.py index 7ddc4d18a78..244cfceb388 100644 --- a/py/torch_tensorrt/dynamo/lowering/passes/pass_utils.py +++ b/py/torch_tensorrt/dynamo/lowering/passes/pass_utils.py @@ -9,6 +9,8 @@ # When True, clean_up_graph_after_modifications only marks the graph dirty and # skips DCE/lint/recompile. post_lowering enables this around the pass list so # multiple cheap FX repairs share one cleanup cycle (Naren: "one iteration"). +# The conversion flush keeps DCE/lint but skips recompile: TRTInterpreter walks +# gm.graph.nodes and never calls gm.forward(). _defer_graph_cleanup: bool = False _graph_cleanup_pending: bool = False @@ -23,20 +25,31 @@ def set_defer_graph_cleanup(enabled: bool) -> None: def flush_deferred_graph_cleanup( gm: torch.fx.GraphModule, + *, + recompile: bool = True, ) -> torch.fx.GraphModule: - """Run a real cleanup if any deferred clean_up call happened.""" + """Run a real cleanup if any deferred clean_up call happened. + + ``recompile`` rebuilds ``gm.forward`` from the graph. Conversion does not + need that; keep it for callers that execute the module as Python. + """ global _graph_cleanup_pending if _graph_cleanup_pending: _graph_cleanup_pending = False - return _run_graph_cleanup(gm) + return _run_graph_cleanup(gm, recompile=recompile) return gm -def _run_graph_cleanup(gm: torch.fx.GraphModule) -> torch.fx.GraphModule: - """Runs dead-code elimination, linting, and recompilation for graph, in-place""" +def _run_graph_cleanup( + gm: torch.fx.GraphModule, + *, + recompile: bool = True, +) -> torch.fx.GraphModule: + """Runs dead-code elimination, linting, and optionally recompilation.""" gm.graph.eliminate_dead_code() gm.graph.lint() - gm.recompile() + if recompile: + gm.recompile() return gm @@ -47,6 +60,7 @@ def clean_up_graph_after_modifications( If deferred cleanup is enabled (see ``set_defer_graph_cleanup``), records that a cleanup is needed and returns immediately so callers can batch mutations. + Non-deferred calls still recompile because those callers may execute Python. """ global _graph_cleanup_pending if _defer_graph_cleanup: diff --git a/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py b/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py index 0dfb74bfc6f..2a321367b85 100644 --- a/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py +++ b/py/torch_tensorrt/dynamo/partitioning/_adjacency_partitioner.py @@ -382,10 +382,11 @@ def partition( Returns: torch.fx.GraphModule, OpSupportTester """ - # Ensure graph is clean prior to partitioning + # Ensure graph is clean prior to partitioning. Skip recompile: AccNodesFinder + # and split() walk graph.nodes. GraphModule construction of split subgraphs + # rebuilds forward() for runtime wrappers and _run_on_gpu fallbacks. gm.graph.eliminate_dead_code() gm.graph.lint() - gm.recompile() # Construct supported_ops = OpSupportTester(torch_executed_ops=torch_executed_ops) diff --git a/tests/py/dynamo/conversion/harness.py b/tests/py/dynamo/conversion/harness.py index 9442fd966c5..62239670f30 100644 --- a/tests/py/dynamo/conversion/harness.py +++ b/tests/py/dynamo/conversion/harness.py @@ -367,7 +367,9 @@ def generate_graph( fx_module = torch.fx.symbolic_trace(mod) if enable_passes: - fx_module = post_lowering(fx_module, settings) + # TRTTestCase.run_test executes the module as Python for the + # reference, so rebuild forward after graph rewrites. + fx_module = post_lowering(fx_module, settings, recompile=True) if propagate_shapes: # TODO: This is currently being used to test embedding_bag_aten due to https://github.com/pytorch/TensorRT/issues/2843 diff --git a/tests/py/dynamo/lowering/test_skip_conversion_recompile.py b/tests/py/dynamo/lowering/test_skip_conversion_recompile.py new file mode 100644 index 00000000000..13805801bcd --- /dev/null +++ b/tests/py/dynamo/lowering/test_skip_conversion_recompile.py @@ -0,0 +1,104 @@ +import unittest +from unittest.mock import patch + +import torch +from torch_tensorrt.dynamo._settings import CompilationSettings +from torch_tensorrt.dynamo.lowering import post_lowering +from torch_tensorrt.dynamo.lowering.passes.pass_utils import ( + clean_up_graph_after_modifications, +) + + +def _count_recompile(fn: object) -> int: + calls = {"n": 0} + orig = torch.fx.GraphModule.recompile + + def counting(self: torch.fx.GraphModule) -> None: + calls["n"] += 1 + return orig(self) + + with patch.object(torch.fx.GraphModule, "recompile", counting): + fn() + return calls["n"] + + +class _AddOne(torch.nn.Module): + def forward(self, value: torch.Tensor) -> torch.Tensor: + return value + 1 + + +def _exported_add_one() -> torch.fx.GraphModule: + return torch.export.export(_AddOne(), (torch.ones(2, 2),)).module() + + +class TestSkipConversionRecompile(unittest.TestCase): + def test_post_lowering_skips_recompile(self) -> None: + gm = _exported_add_one() + n = _count_recompile(lambda: post_lowering(gm, CompilationSettings())) + self.assertEqual(n, 0) + + def test_post_lowering_still_eliminates_dead_code(self) -> None: + gm = _exported_add_one() + output = next(n for n in reversed(list(gm.graph.nodes)) if n.op == "output") + inp = next(n for n in gm.graph.nodes if n.op == "placeholder") + with gm.graph.inserting_before(output): + gm.graph.call_function(torch.ops.aten.mul.Tensor, args=(inp, inp)) + mul_before = sum( + 1 for n in gm.graph.nodes if n.target is torch.ops.aten.mul.Tensor + ) + self.assertEqual(mul_before, 1) + + post_lowering(gm, CompilationSettings()) + + mul_after = sum( + 1 for n in gm.graph.nodes if n.target is torch.ops.aten.mul.Tensor + ) + self.assertEqual(mul_after, 0) + + def test_interpreter_runs_without_recompile(self) -> None: + value = torch.ones(2, 2) + gm = torch.export.export(_AddOne(), (value,)).module() + gm = post_lowering(gm, CompilationSettings()) + out = torch.fx.Interpreter(gm).run(value) + torch.testing.assert_close(out, value + 1) + + def test_clean_up_outside_defer_still_recompiles(self) -> None: + gm = _exported_add_one() + n = _count_recompile(lambda: clean_up_graph_after_modifications(gm)) + self.assertGreater(n, 0) + + def test_post_lowering_recompile_true_rebuilds_forward(self) -> None: + value = torch.ones(2, 2) + gm = torch.export.export(_AddOne(), (value,)).module() + n = _count_recompile( + lambda: post_lowering(gm, CompilationSettings(), recompile=True) + ) + self.assertGreater(n, 0) + torch.testing.assert_close(gm(value), value + 1) + + def test_fast_partition_does_not_recompile_input_module(self) -> None: + from torch_tensorrt.dynamo.partitioning import fast_partition + + gm = _exported_add_one() + calls = {"n": 0} + orig = gm.recompile + + def counting(*args: object, **kwargs: object) -> None: + calls["n"] += 1 + return orig(*args, **kwargs) + + gm.recompile = counting # type: ignore[method-assign] + partitioned, _ = fast_partition( + gm, + min_block_size=1, + require_full_compilation=True, + assume_full_support=True, + skip_fusion=True, + ) + self.assertEqual(calls["n"], 0) + value = torch.ones(2, 2) + torch.testing.assert_close(partitioned(value), value + 1) + + +if __name__ == "__main__": + unittest.main() diff --git a/tools/llm/tests/test_static_cache.py b/tools/llm/tests/test_static_cache.py index 2683b63a009..714bed55ad6 100644 --- a/tools/llm/tests/test_static_cache.py +++ b/tools/llm/tests/test_static_cache.py @@ -256,7 +256,7 @@ def transform_gm_with_kv_cache(exported_program: torch.export.ExportedProgram, a exported_program = exported_program.run_decompositions(get_decompositions(False)) gm = exported_program.module() - gm = post_lowering(gm, settings) + gm = post_lowering(gm, settings, recompile=True) return gm