From 7c7cbb6ed461759e3b8876220d56f6250401369e Mon Sep 17 00:00:00 2001 From: Anthony Shoumikhin Date: Tue, 25 Aug 2026 17:52:30 -0700 Subject: [PATCH] fix(dynamo): keep the complex decomposition inside the TRT op set The complex lowering pass hands the interior expansion to PyTorch's decompose_complex_in_graph, which re-traces the module with make_fx. The module that comes back is therefore built from whatever the complex tensor subclass dispatched to, not from the ops that survived the earlier stages of the same pass list, and three kinds of leftovers reach the partitioner: * aten.stack and aten.alias, which the decomposition table run at export time expands and drops, are back; * constants that constant folding had already reduced to one frozen tensor are live chains of ops again, and the converters for those ops want a TRT tensor rather than a weight; * aten.view_as_real and aten.view_as_complex are left at the seams. The last one only shows up when the complex region is interior, that is when the graph takes real tensors and returns real tensors, as a rotary embedding does. The boundary normalizer only knew about aten.complex, aten.real and aten.imag, which is what a graph with complex inputs or outputs meets upstream at. None of these have a TensorRT converter, so the partitioner cut the graph around them: a rotary attention module that used to build one engine came back as three, with PyTorch blocks in between. Pass the decomposition table to the retrace, fold constants again after it, and teach the normalizer the two view_as_real seams. The rotary attention module builds a single engine again and matches eager output. --- .../passes/complex_decomposition_adapter.py | 126 ++++++++++++++---- .../test_complex_decomposition_adapter.py | 125 ++++++++++++++++- 2 files changed, 224 insertions(+), 27 deletions(-) diff --git a/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py b/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py index c47fd31860..05a6f4f2ff 100644 --- a/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py +++ b/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py @@ -8,9 +8,13 @@ 1. capture the complex I/O signature before rewriting (drives the boundary adapters in ``_compiler._insert_complex_io_adapters``); 2. call upstream to expand every complex op into real ops on separate re/im; - 3. normalize the SoA seams (``aten.complex`` / ``aten.real`` / ``aten.imag``) - into the interleaved ``[..., 2]`` layout the rest of the TRT flow expects, - so the engine only ever sees real tensors (Option A in the RFC). + 3. put back what the retrace undid: upstream returns a module rebuilt from + scratch, so the decomposition and constant-folding stages that ran earlier + find their work partly reverted; + 4. normalize the seams (``aten.view_as_real`` / ``aten.complex`` / + ``aten.real`` / ``aten.imag``) into the interleaved ``[..., 2]`` layout the + rest of the TRT flow expects, so the engine only ever sees real tensors + (Option A in the RFC). Gated by ``settings.use_complex_decomposition``; falls back to the legacy pass on older torch (the upstream API only exists in torch>=2.14.0.dev), or if @@ -26,6 +30,10 @@ from torch.fx import GraphModule from torch_tensorrt._features import has_complex_decomposition from torch_tensorrt.dynamo._settings import CompilationSettings +from torch_tensorrt.dynamo.lowering.passes.constant_folding import constant_fold +from torch_tensorrt.dynamo.lowering.passes.mark_constant_fold_exclusions import ( + mark_constant_fold_exclusions, +) from torch_tensorrt.dynamo.lowering.passes.pass_utils import ( clean_up_graph_after_modifications, ) @@ -91,7 +99,9 @@ def complex_decomposition_adapter( # call either fully succeeds or raises; it never partially mutates # anything, since it builds a new GraphModule rather than editing # the existing one in place). - decomposed_gm = decompose_complex_in_graph(gm, flat_args) + decomposed_gm = decompose_complex_in_graph( + gm, flat_args, decompositions=_trt_decomposition_table(settings) + ) except Exception as e: # decompose_complex_in_graph is upstream, experimental PyTorch code # (torch._functorch._aot_autograd.complex_decomposition) with @@ -107,11 +117,18 @@ def complex_decomposition_adapter( return complex_graph_detection(gm, settings) gm = decomposed_gm - # (3) Normalize SoA seams into the interleaved [..., 2] layout used by TRT. + # (3) Fold constants again. The retrace builds the module from scratch, so + # a constant that the constant-folding stage had already reduced to one + # frozen tensor comes back as a live chain of ops reading it, and the + # converters for those ops expect a TRT tensor rather than a weight. + gm = mark_constant_fold_exclusions(gm, settings) + gm = constant_fold(gm, settings) + + # (4) Normalize the seams into the interleaved [..., 2] layout used by TRT. gm = _normalize_complex_boundary_for_trt(gm) gm = clean_up_graph_after_modifications(gm) - # (4) Re-attach the captured I/O signature onto the RETURNED module so + # (5) Re-attach the captured I/O signature onto the RETURNED module so # _insert_complex_io_adapters can restore the complex boundary. gm.meta["complex_output_indices"] = complex_output_indices gm.meta["complex_input_names"] = complex_input_names @@ -123,6 +140,29 @@ def complex_decomposition_adapter( return gm +def _trt_decomposition_table( + settings: CompilationSettings, +) -> dict[Any, Any]: + """The op set the rest of the TRT flow expects to see. + + ``decompose_complex_in_graph`` re-traces through ``make_fx``, so the module + it returns is built from whatever ``ComplexTensor`` dispatched to, not from + the ops that survived the decomposition run at export time. Without a table + that retrace re-introduces ops the flow has already decided to expand + (``aten.stack``) or to drop (``aten.alias``), long after the stage that + would have handled them, and the partitioner then cuts the graph around + them. Handing the retrace the same table keeps the two in step. + """ + from torch_tensorrt.dynamo.lowering import get_decompositions + + return get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + settings.use_distributed_mode_trace, + use_fp32_acc=settings.use_fp32_acc, + ) + + def _graph_has_complex(gm: GraphModule) -> bool: from torch_tensorrt.dynamo.utils import COMPLEX_DTYPES @@ -174,20 +214,30 @@ def _fake_flat_args(gm: GraphModule) -> list[torch.Tensor]: def _normalize_complex_boundary_for_trt(gm: GraphModule) -> GraphModule: - """Fold upstream's SoA seams into the interleaved ``[..., 2]`` real layout. + """Fold upstream's seams into the interleaved ``[..., 2]`` real layout. - Upstream leaves the graph in terms of separate re/im joined by - ``aten.complex(re, im)`` and unpacked by ``aten.real`` / ``aten.imag``. TRT - has no converter for any of those. We: + Every op below reinterprets a complex value as a pair of real halves, or + the other way round, and TRT has no converter for any of them. Leaving one + in the graph does not just fail an op check: the partitioner cuts the graph + around it and the model comes back split into several engines with PyTorch + blocks in between. We rewrite: + * ``view_as_real(z)`` -> the interleaved ``[..., 2]`` tensor it would + have returned, for each producer ``z`` can have (see Pass A) * ``real(z)`` / ``imag(z)`` where ``z = complex(re, im)`` -> re / im (cancel the pack/unpack round-trip) - * remaining ``aten.complex(re, im)`` -> ``stack([re, im], -1)`` tagged as + * remaining ``aten.complex(re, im)`` -> ``cat([re, im], -1)`` tagged as complex-layout, so the [..., 2] tensor flows to the boundary adapters. - Confirmed against real decompose_complex_in_graph output: the boundary op - set is aten.complex/real/imag as assumed above, unpacked via real/imag - (not view_as_real). + Which of these appear depends on where the complex region sits. A graph + whose own inputs or outputs are complex meets upstream at ``aten.complex``, + because that is what ``ComplexTensor`` packs with on the way out. A graph + that is real on both ends but complex in the middle, as a rotary embedding + is, meets it at ``view_as_real`` instead: the entry ``view_as_complex`` is + left alone (its argument is a plain real tensor, so ``ComplexTensor`` never + sees it), and ``ComplexTensor`` splits its halves with ``torch.real`` / + ``torch.imag``, which are composite ops that expand to ``view_as_real`` + plus ``select`` before any subclass can intercept them. """ g = gm.graph aten = torch.ops.aten @@ -210,7 +260,43 @@ def _propagate_meta(node: torch.fx.Node) -> None: with fake_mode: node.meta["val"] = node.target(*arg_vals) - # Pass A: cancel real(complex(re,im)) / imag(complex(re,im)) round-trips. + def _pack_interleaved(re: Any, im: Any) -> torch.fx.Node: + """Build the ``[..., 2]`` tensor holding ``re`` and ``im`` side by side. + + Callers set the insertion point. ``cat`` of two unsqueezed halves + rather than ``stack`` because the stage that expands ``stack`` has + already run by the time this pass does. + """ + re_u = g.call_function(aten.unsqueeze.default, (re, -1)) + _propagate_meta(re_u) + im_u = g.call_function(aten.unsqueeze.default, (im, -1)) + _propagate_meta(im_u) + packed = g.call_function(aten.cat.default, ([re_u, im_u], -1)) + _propagate_meta(packed) + return packed + + # Pass A: rewrite view_as_real over each producer it can have. + for node in list(g.nodes): + if node.op != "call_function" or node.target != aten.view_as_real.default: + continue + src = node.args[0] + if not isinstance(src, torch.fx.Node): + continue + if src.target == aten.view_as_complex.default: + # view_as_complex consumes a trailing dimension of 2 and + # view_as_real puts it back holding the same values, so the pair is + # an identity on the real tensor that went in. Dropping it leaves + # the halves to be picked out of that tensor directly. + replacement = src.args[0] + elif src.target == aten.complex.default: + with g.inserting_before(node): + replacement = _pack_interleaved(src.args[0], src.args[1]) + else: + continue + node.replace_all_uses_with(replacement) + g.erase_node(node) + + # Pass B: cancel real(complex(re,im)) / imag(complex(re,im)) round-trips. for node in list(g.nodes): if node.op != "call_function" or node.target not in ( aten.real.default, @@ -233,18 +319,12 @@ def _propagate_meta(node: torch.fx.Node) -> None: node.replace_all_uses_with(src.args[idx]) g.erase_node(node) - # Pass B: turn surviving aten.complex(re, im) into stack([re, im], -1). + # Pass C: turn surviving aten.complex(re, im) into the interleaved layout. for node in list(g.nodes): if node.op != "call_function" or node.target != aten.complex.default: continue - re, im = node.args[0], node.args[1] with g.inserting_before(node): - re_u = g.call_function(aten.unsqueeze.default, (re, -1)) - _propagate_meta(re_u) - im_u = g.call_function(aten.unsqueeze.default, (im, -1)) - _propagate_meta(im_u) - packed = g.call_function(aten.cat.default, ([re_u, im_u], -1)) - _propagate_meta(packed) + packed = _pack_interleaved(node.args[0], node.args[1]) packed.meta["is_complex_layout"] = True node.replace_all_uses_with(packed) g.erase_node(node) diff --git a/tests/py/dynamo/lowering/test_complex_decomposition_adapter.py b/tests/py/dynamo/lowering/test_complex_decomposition_adapter.py index 8052a1fb72..0eb4aca2b3 100644 --- a/tests/py/dynamo/lowering/test_complex_decomposition_adapter.py +++ b/tests/py/dynamo/lowering/test_complex_decomposition_adapter.py @@ -3,9 +3,10 @@ These tests exercise the TRT-specific glue around PyTorch's upstream complex decomposition -- they do NOT require a GPU or a TRT build: - * _normalize_complex_boundary_for_trt: folds the aten.complex / real / imag - seams that decompose_complex_in_graph leaves behind into the interleaved - [..., 2] real layout the rest of the TRT flow expects. + * _normalize_complex_boundary_for_trt: folds the aten.view_as_real / + aten.complex / real / imag seams that decompose_complex_in_graph leaves + behind into the interleaved [..., 2] real layout the rest of the TRT flow + expects. * the torch-version feature gate: when the upstream API is unavailable the adapter must fall back to the legacy complex_graph_detection pass. """ @@ -81,6 +82,118 @@ def forward(self, re, im): assert torch.equal(out[..., 1], im) +def test_normalize_folds_view_as_real_of_view_as_complex(): + """view_as_real(view_as_complex(x)) -> x, and both nodes are erased.""" + + class M(torch.nn.Module): + def forward(self, x): + z = torch.ops.aten.view_as_complex.default(x) + return torch.ops.aten.view_as_real.default(z) + + from torch_tensorrt.dynamo.lowering.passes.pass_utils import ( + clean_up_graph_after_modifications, + ) + + x = torch.randn(3, 2) + gm = torch.fx.symbolic_trace(M()) + + gm = cda._normalize_complex_boundary_for_trt(gm) + # The view_as_complex node is left without users, and the adapter runs this + # right after the pass to collect nodes in exactly that state. + gm = clean_up_graph_after_modifications(gm) + targets = _targets(gm) + + assert aten.view_as_real.default not in targets + assert aten.view_as_complex.default not in targets + assert torch.equal(gm(x), x) + + +def test_normalize_folds_view_as_real_of_complex(): + """view_as_real(complex(re,im)) -> the interleaved [..., 2] layout. + + Without this, Pass C would rewrite the complex() node into that layout and + leave a view_as_real reading it, which is no longer a complex tensor. + """ + + class M(torch.nn.Module): + def forward(self, re, im): + z = torch.ops.aten.complex.default(re, im) + return torch.ops.aten.view_as_real.default(z) + + re = torch.randn(3) + im = torch.randn(3) + gm = torch.fx.symbolic_trace(M()) + + gm = cda._normalize_complex_boundary_for_trt(gm) + targets = _targets(gm) + + assert aten.view_as_real.default not in targets + assert aten.complex.default not in targets + + out = gm(re, im) + assert out.shape == (3, 2) + assert torch.equal(out[..., 0], re) + assert torch.equal(out[..., 1], im) + + +# --------------------------------------------------------------------------- +# whole pass over a graph that is real at both ends and complex in the middle +# --------------------------------------------------------------------------- + + +@pytest.mark.skipif( + not cda.has_complex_decomposition(), + reason="decompose_complex_in_graph requires torch>=2.14.dev", +) +def test_interior_complex_region_leaves_nothing_trt_cannot_convert(): + """A rotary-embedding shaped graph must come out with real ops only. + + The complex region here sits in the middle: the graph takes a real tensor + and returns one, so upstream meets it at view_as_real rather than at + aten.complex, and its retrace also brings back ops the decomposition and + constant-folding stages had already dealt with. Any of those left behind + is not just an unconverted op: the partitioner cuts the graph around it and + the model comes back as several engines with PyTorch blocks in between. + """ + from torch_tensorrt.dynamo._settings import CompilationSettings + from torch_tensorrt.dynamo.lowering import get_decompositions + + class M(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer( + "freqs", + torch.complex(torch.randn(1, 2, 1, 1), torch.randn(1, 2, 1, 1)), + ) + + def forward(self, x): + z = torch.view_as_complex(x.reshape(1, 2, 1, 1, 2)) + return torch.view_as_real(z * self.freqs) + + model = M().eval() + x = torch.randn(1, 2, 1, 2) + expected = model(x) + + gm = ( + torch.export.export(model, (x,)) + .run_decompositions(get_decompositions()) + .module() + ) + out = cda.complex_decomposition_adapter(gm, CompilationSettings()) + + targets = set(_targets(out)) + for target in ( + aten.view_as_complex.default, + aten.view_as_real.default, + aten.complex.default, + aten.alias.default, + aten.stack.default, + ): + assert target not in targets, f"{target} survived the pass" + + torch.testing.assert_close(out(x), expected) + + # --------------------------------------------------------------------------- # feature gate / fallback # --------------------------------------------------------------------------- @@ -172,7 +285,11 @@ def fake_decompose(g, flat_args, *a, **k): ) monkeypatch.setattr(cda, "has_complex_decomposition", lambda: True) - out = cda.complex_decomposition_adapter(gm, settings=object()) + # Real settings, not a stub: the adapter reads the decomposition options off + # them to build the table it hands to the retrace. + from torch_tensorrt.dynamo._settings import CompilationSettings + + out = cda.complex_decomposition_adapter(gm, settings=CompilationSettings()) # Functional: we did NOT get the original module back... assert out is not gm