From b2ebb04f681e2cce25d163c41a28a4fecfc90103 Mon Sep 17 00:00:00 2001 From: Andrew Pullin Date: Mon, 21 Sep 2026 23:26:37 -0700 Subject: [PATCH] Skip redundant run_decompositions when no ops match decomp table (#18496) Summary: `_gen_edge_manager_for_partitioners` can call `program.run_decompositions(table)` up to three times. Each call re-exports the program through `make_fx`, retracing every node through FakeTensor dispatch, even when a previous pass already removed every operator covered by the next decomposition table. Before each nonempty-table replay, scan the root and every descendant `GraphModule`. Skip `run_decompositions` when no `call_function` target matches the table. Empty tables still run to preserve functionalization. Overload packets are checked against their constituent overloads, and nested-region graphs with their own decomposition policy conservatively force a replay. This keeps the optimization correct for `cond`, `map`, `scan`, `while_loop`, and `invoke_subgraph` bodies rather than limiting the scan to the previously enumerated control-flow operators. ## Benchmark Synthetic calibration lowering suite, five models: Comparison revision: 79.000 s / 79.516 s This change: 65.637 s / 65.512 s Delta: -17.3% mean / -17.6% warm CombinedControl Ethos-U55, structured `LOWERING.duration_ms`: Comparison revision: 145.801 s / 145.888 s This change: 133.949 s / 132.096 s Delta: -8.8% mean / -9.5% warm Differential Revision: D96489903 --- exir/program/_program.py | 63 +++++++++-- exir/program/test/test_program.py | 181 ++++++++++++++++++++++++++++++ 2 files changed, 235 insertions(+), 9 deletions(-) diff --git a/exir/program/_program.py b/exir/program/_program.py index 94c5cad4786..7032e9e5208 100644 --- a/exir/program/_program.py +++ b/exir/program/_program.py @@ -33,10 +33,7 @@ from executorch.exir.emit import emit_program, EmitterOutput from executorch.exir.emit._emitter import _DelegateDebugIdentifierMap from executorch.exir.error import ExportError -from executorch.exir.graph_module import ( - contains_any_call_fn_target_op, - get_control_flow_submodules, -) +from executorch.exir.graph_module import get_control_flow_submodules from executorch.exir.operator.convert import _pybind_schema_to_native_schema from executorch.exir.operator.util import _QUANT_PRIMITIVES from executorch.exir.pass_base import PassBase @@ -1129,7 +1126,51 @@ def _apply_pre_decomposition_transforms( return program -def _gen_edge_manager_for_partitioners( +def _has_decomposable_ops( + program: "ExportedProgram", + decomp_table: dict, +) -> bool: + """Check if any ops in the program match the decomposition table. + + Nested GraphModules include control-flow and invoke_subgraph bodies that + run_decompositions also retraces. Returns True for empty tables because + that is the functionalization-only path. + """ + if not decomp_table: + return True + + for graph_module in program.graph_module.modules(): + if not isinstance(graph_module, torch.fx.GraphModule): + continue + nested_config = graph_module.meta.get("nested_region_config") + if ( + nested_config is not None + and getattr(nested_config, "decompositions", None) is not None + ): + return True + for node in graph_module.graph.nodes: + if node.op != "call_function": + continue + raw_target = node.target + target = raw_target + if not isinstance(target, torch._ops.OpOverload): + target = getattr(target, "_op", target) + if target in decomp_table: + return True + packet = ( + raw_target + if isinstance(raw_target, torch._ops.OpOverloadPacket) + else getattr(raw_target, "_op", None) + ) + if isinstance(packet, torch._ops.OpOverloadPacket) and any( + getattr(packet, overload) in decomp_table + for overload in packet.overloads() + ): + return True + return False + + +def _gen_edge_manager_for_partitioners( # noqa: C901 partitioner: Dict[str, List[Partitioner]], aten_programs: Dict[str, ExportedProgram], config: EdgeCompileConfig, @@ -1166,7 +1207,8 @@ def _gen_edge_manager_for_partitioners( table = _default_decomposition_table() for op in config.preserve_ops: table.pop(op, None) - program = program.run_decompositions(table) + if _has_decomposable_ops(program, table): + program = program.run_decompositions(table) # Process each partitioner individually using their specific requirements for curr_partitioner in partitioners_for_program: @@ -1190,7 +1232,7 @@ def _gen_edge_manager_for_partitioners( # functionalized the graph. This second call only applies # operator decompositions from the remaining table, so it is # safe to skip when none of those targets occur in the graph. - if contains_any_call_fn_target_op(program.graph_module, table): + if _has_decomposable_ops(program, table): program = program.run_decompositions(table) final_ops_to_preserve.update(ops_needing_preservation) else: @@ -1205,7 +1247,8 @@ def _gen_edge_manager_for_partitioners( table.pop(op, None) # First pass of decompositions with this partitioner's preserved ops - program = program.run_decompositions(table) + if _has_decomposable_ops(program, table): + program = program.run_decompositions(table) # Filter ops using EDGE_DO_NOT_DECOMP temp_partitioner_dict = {name: [curr_partitioner]} @@ -1218,7 +1261,9 @@ def _gen_edge_manager_for_partitioners( final_ops_to_preserve.update(preserved_ops) # Second pass of decompositions with this partitioner's preserved ops after filtering - program = program.run_decompositions(_default_decomposition_table()) + full_table = _default_decomposition_table() + if _has_decomposable_ops(program, full_table): + program = program.run_decompositions(full_table) # Restore ops from edge_no_decomp_namespace to aten ops _restore_transformed_ops_to_aten_ops(program) diff --git a/exir/program/test/test_program.py b/exir/program/test/test_program.py index 4f6ef27b3a7..455ea244075 100644 --- a/exir/program/test/test_program.py +++ b/exir/program/test/test_program.py @@ -22,6 +22,7 @@ from executorch.exir.pass_base import ExportPass from executorch.exir.passes import MemoryPlanningPass from executorch.exir.program._program import ( + _has_decomposable_ops, _transform, EdgeProgramManager, ExecutorchProgramManager, @@ -950,3 +951,183 @@ def __init__(self): ) self.assertTrue(issubclass(transformed.verifiers[0], MyVerifier)) self.assertFalse(issubclass(program.verifiers[0], MyVerifier)) + + +def _placeholder_decomp(*args: Any, **kwargs: Any) -> None: + # _has_decomposable_ops only checks table membership, never calls the value. + return None + + +def _call_function_targets(gm: torch.fx.GraphModule) -> list[Any]: + return [n.target for n in gm.graph.nodes if n.op == "call_function"] + + +def _all_call_function_targets(gm: torch.fx.GraphModule) -> list[Any]: + return [ + node.target + for module in gm.modules() + if isinstance(module, torch.fx.GraphModule) + for node in module.graph.nodes + if node.op == "call_function" + ] + + +def _all_graph_module_code(gm: torch.fx.GraphModule) -> dict[str, str]: + return { + name: module.code + for name, module in gm.named_modules() + if isinstance(module, torch.fx.GraphModule) + } + + +class HasDecomposableOpsTest(unittest.TestCase): + class AddMul(torch.nn.Module): + def forward(self, x: torch.Tensor) -> torch.Tensor: + return x * 2 + 3 + + class Cond(torch.nn.Module): + def forward(self, pred: torch.Tensor, x: torch.Tensor) -> torch.Tensor: + def true_fn(x: torch.Tensor) -> torch.Tensor: + return x.sin() + + def false_fn(x: torch.Tensor) -> torch.Tensor: + return x.cos() + + return torch.cond(pred, true_fn, false_fn, [x]) + + class WhileLinear(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.linear = torch.nn.Linear(2, 2) + + def forward( + self, iterations: torch.Tensor, x: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + def cond_fn(it: torch.Tensor, value: torch.Tensor) -> torch.Tensor: + return it > 0 + + def body_fn( + it: torch.Tensor, value: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + return it - 1, self.linear(value) + + return torch._higher_order_ops.while_loop(cond_fn, body_fn, (iterations, x)) + + def _assert_false_guard_matches_noop( + self, + module: torch.nn.Module, + inputs: tuple[Any, ...], + ) -> None: + program = _export(module, inputs, pre_dispatch=True).run_decompositions({}) + table = _default_decomposition_table() + for target in _call_function_targets(program.graph_module): + if target in table: + table.pop(target) + + self.assertFalse(_has_decomposable_ops(program, table)) + + decomposed = copy.deepcopy(program).run_decompositions(table) + self.assertEqual( + _all_graph_module_code(program.graph_module), + _all_graph_module_code(decomposed.graph_module), + ) + self.assertEqual(program.graph_signature, decomposed.graph_signature) + torch.testing.assert_close( + program.module()(*inputs), + decomposed.module()(*inputs), + ) + + def test_empty_table_returns_true(self) -> None: + # Empty table is the functionalize-only path, which the guard never skips. + ep = export(self.AddMul(), (torch.randn(4),), strict=True) + self.assertTrue(_has_decomposable_ops(ep, {})) + + def test_returns_true_when_graph_has_matching_op(self) -> None: + ep = export(self.AddMul(), (torch.randn(4),), strict=True) + a_target = _call_function_targets(ep.graph_module)[0] + self.assertTrue(_has_decomposable_ops(ep, {a_target: _placeholder_decomp})) + + def test_packet_target_matches_overload_decomposition(self) -> None: + ep = export(self.AddMul(), (torch.randn(4),), strict=True) + add_node = next( + node + for node in ep.graph_module.graph.nodes + if node.target == torch.ops.aten.add.Tensor + ) + add_node.target = torch.ops.aten.add + self.assertIs(add_node.target, torch.ops.aten.add) + + self.assertTrue( + _has_decomposable_ops( + ep, + {torch.ops.aten.add.Tensor: _placeholder_decomp}, + ) + ) + + def test_nested_region_local_decompositions_force_run(self) -> None: + class NestedConfig: + decompositions: dict[Any, Any] = {} + + ep = export(self.AddMul(), (torch.randn(4),), strict=True) + ep.graph_module.meta["nested_region_config"] = NestedConfig() + + self.assertTrue( + _has_decomposable_ops( + ep, + {torch.ops.aten.sin.default: _placeholder_decomp}, + ) + ) + + def test_returns_false_when_no_matching_op(self) -> None: + ep = export(self.AddMul(), (torch.randn(4),), strict=True) + # Start from the real table and drop every op the graph actually uses, + # so no remaining key can match. + table = _default_decomposition_table() + for target in _call_function_targets(ep.graph_module): + table.pop(target, None) + self.assertFalse(_has_decomposable_ops(ep, table)) + + def test_detects_op_in_control_flow_submodule(self) -> None: + # The matching op (sin) lives only inside the cond branches, not the + # top-level graph -- this exercises the nested GraphModule traversal + # that a top-level-only scan would miss. + ep = export(self.Cond(), (torch.tensor(True), torch.randn(3)), strict=True) + self.assertNotIn( + torch.ops.aten.sin.default, _call_function_targets(ep.graph_module) + ) + table = {torch.ops.aten.sin.default: _placeholder_decomp} + self.assertTrue(_has_decomposable_ops(ep, table)) + + def test_detects_op_in_while_loop_submodule(self) -> None: + linear = torch.ops.aten.linear.default + program = export( + self.WhileLinear(), + (torch.tensor(3), torch.randn(2, 2)), + strict=True, + ).run_decompositions({}) + table = _default_decomposition_table() + linear_table = {linear: table[linear]} + + self.assertNotIn(linear, _call_function_targets(program.graph_module)) + self.assertIn(linear, _all_call_function_targets(program.graph_module)) + self.assertTrue(_has_decomposable_ops(program, linear_table)) + + decomposed = copy.deepcopy(program).run_decompositions(linear_table) + self.assertNotIn( + linear, + _all_call_function_targets(decomposed.graph_module), + ) + + def test_false_implies_run_decompositions_is_noop(self) -> None: + self._assert_false_guard_matches_noop( + self.AddMul(), + (torch.randn(4),), + ) + + def test_false_guard_matches_noop_for_complex_models(self) -> None: + for model in (TestLSTM(), TestLinearSDPACombined()): + with self.subTest(model=type(model).__name__): + self._assert_false_guard_matches_noop( + model, + model._get_random_inputs(), + )