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
17 changes: 13 additions & 4 deletions py/torch_tensorrt/dynamo/lowering/passes/_aten_lowering_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,

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.

Check out torch.fx._lazy_graph_module._LazyGraphModule.

AI:
Defaulting recompile=False assumes the lowered module is never run as Python. But compile_module returns it unconverted when a graph has fewer supported ops than min_block_size, and the torch.compile backend then runs its stale forward.

backends.py:367 calls post_lowering(gm, settings) with the default, then compile_module. _compiler.py:1381 returns gm when num_supported_ops < min_block_size, which happens often with small Dynamo graph-break fragments. That gm.forward was generated before lowering. It still runs pre-lowering code, and it can refer to attributes that constant_fold deleted.

This PR already had to pass recompile=True in harness.py and test_static_cache.py for this reason. The same applies to compile(..., dryrun=True) at :1415. There's no test for the early-return path. I found this by inspection and haven't run it.

The fix is to convert to torch.fx._lazy_graph_module._LazyGraphModule instead of adding the flag. Its recompile() only marks the code dirty, and code generation runs on first access to forward. I measured 56.75 ms for GraphModule.recompile() against 0.01 ms for the lazy version on a 2000-node graph. That's the same saving, with no stale-module risk.

) -> 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:
Expand Down
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
24 changes: 19 additions & 5 deletions py/torch_tensorrt/dynamo/lowering/passes/pass_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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


Expand All @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
4 changes: 3 additions & 1 deletion tests/py/dynamo/conversion/harness.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
104 changes: 104 additions & 0 deletions tests/py/dynamo/lowering/test_skip_conversion_recompile.py
Original file line number Diff line number Diff line change
@@ -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()
2 changes: 1 addition & 1 deletion tools/llm/tests/test_static_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading