From 5d9ab465a9fdff7c8d4775957c2535d3dd2d711f Mon Sep 17 00:00:00 2001 From: Yufeng Shi Date: Fri, 17 Jul 2026 16:00:33 +0100 Subject: [PATCH] Arm backend: Rerun duplicate-user fusion after TOSA lowering Late TOSA transformations can introduce equivalent operations after the first FuseDuplicateUsersPass invocation. Rerun it after TOSA and shape transformations, before output nodes are made unique. TOSA-FP operator comparisons: | Model | Before | After | Reduction | |--------------------|-------:|------:|----------:| | SD3 | 1,663 | 1,573 | 90 (5.4%) | | InceptionV3 | 762 | 746 | 16 (2.1%) | | Conformer delegate | 537 | 494 | 43 (8.0%) | All three reference-output tests pass with late fusion enabled. Add regression coverage for the late fusion and output uniqueness pass ordering. Change-Id: Ia4e2335e05b18d16b93c7e035a5502b8f050855c Signed-off-by: Yufeng Shi --- backends/arm/_passes/arm_pass_manager.py | 7 ++- .../passes/test_fuse_duplicate_users_pass.py | 46 +++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) diff --git a/backends/arm/_passes/arm_pass_manager.py b/backends/arm/_passes/arm_pass_manager.py index ebedc610864..e508c03d63a 100644 --- a/backends/arm/_passes/arm_pass_manager.py +++ b/backends/arm/_passes/arm_pass_manager.py @@ -655,9 +655,14 @@ def _tosa_pipeline( SymbolicToTosaShapesPass(), InsertDynamicPaddingPass(), FuseConsecutiveConcatShapesPass(), - EnsureUniqueOutputNodesPass(), + # No-op removal can expose duplicate users and outputs, so run + # FuseDuplicateUsersPass and EnsureUniqueOutputNodesPass afterward. RemoveNoopPass(), InsertRescalePass(), + # Late TOSA transformations can introduce duplicate users after + # the first FuseDuplicateUsersPass invocation. + FuseDuplicateUsersPass(), + EnsureUniqueOutputNodesPass(), ] ) diff --git a/backends/arm/test/passes/test_fuse_duplicate_users_pass.py b/backends/arm/test/passes/test_fuse_duplicate_users_pass.py index 027fb6a7919..736ac1a55fc 100644 --- a/backends/arm/test/passes/test_fuse_duplicate_users_pass.py +++ b/backends/arm/test/passes/test_fuse_duplicate_users_pass.py @@ -8,13 +8,17 @@ import executorch.backends.arm.tosa.dialect # noqa: F401 import torch from executorch.backends.arm._passes import FuseDuplicateUsersPass +from executorch.backends.arm._passes.arm_pass_manager import ArmPassManager from executorch.backends.arm.test import common from executorch.backends.arm.test.tester.test_pipeline import PassPipeline +from executorch.backends.arm.tosa.compile_spec import TosaCompileSpec from executorch.backends.arm.tosa.specification import ( TosaLoweringContext, TosaSpecification, ) +from executorch.exir import EdgeCompileConfig, to_edge from executorch.exir.dialects._ops import ops as exir_ops +from torch.export import export from torch.fx import Graph, GraphModule input_t = Tuple[torch.Tensor] # Input x @@ -167,3 +171,45 @@ def test_fuse_duplicate_users_removes_identical_rescale_users(): assert len(rescale_nodes) == 1 output_node = result.graph_module.graph.output_node() assert output_node.args[0] == (rescale_nodes[0], rescale_nodes[0]) + + +class LateDuplicateUsers(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer("first", torch.ones(2, 3)) + self.register_buffer("second", torch.ones(2, 3)) + + def forward(self, x): + return x + self.first, x + self.second + + +def test_fuse_duplicate_users_runs_after_tosa_transformations(): + exported_program = export(LateDuplicateUsers(), (torch.ones(2, 3),), strict=True) + edge_program = to_edge( + exported_program, + compile_config=EdgeCompileConfig(_check_ir_validity=False), + ) + edge_exported_program = edge_program.exported_program() + + graph_module = ArmPassManager( + TosaCompileSpec("TOSA-1.0+FP") + ).transform_to_backend_pipeline( + edge_exported_program, edge_exported_program.graph_module + ) + + add_nodes = [ + node + for node in graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.ADD.default + ] + identity_nodes = [ + node + for node in graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.IDENTITY.default + ] + + graph_module.graph.lint() + assert len(add_nodes) == 1 + assert len(identity_nodes) == 2 + assert all(node.args[0] is add_nodes[0] for node in identity_nodes) + assert graph_module.graph.output_node().args[0] == tuple(identity_nodes)