From 12ef2f47f62566d836af239388290c3d44673bf8 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. Run InsertRescalePass after the final FuseDuplicateUsersPass. Fusing generated RESCALE users can merge distinct quantized paths and produce incorrect results. TOSA-FP operator comparisons: | Model | Before | After | Reduction | |--------------------|-------:|------:|----------:| | SD3 | 1,591 | 1,527 | 64 (4.0%) | | Conformer delegate | 544 | 498 | 46 (8.5%) | Reference-output tests pass with late fusion enabled. Add regression coverage for late fusion, rescale insertion, and output uniqueness pass ordering. Change-Id: Ia4e2335e05b18d16b93c7e035a5502b8f050855c Signed-off-by: Yufeng Shi --- backends/arm/_passes/arm_pass_manager.py | 5 +- .../passes/test_fuse_duplicate_users_pass.py | 64 ++++++++++++++++++- 2 files changed, 67 insertions(+), 2 deletions(-) diff --git a/backends/arm/_passes/arm_pass_manager.py b/backends/arm/_passes/arm_pass_manager.py index da31e6e82d0..522fcc9c8d7 100644 --- a/backends/arm/_passes/arm_pass_manager.py +++ b/backends/arm/_passes/arm_pass_manager.py @@ -667,9 +667,12 @@ def _tosa_pipeline( SymbolicToTosaShapesPass(), InsertDynamicPaddingPass(), FuseConsecutiveConcatShapesPass(), - EnsureUniqueOutputNodesPass(), RemoveNoopPass(), + # Fuse duplicates exposed by late rewrites before inserting rescales; + # fusing generated RESCALE users can corrupt distinct quantized paths. + FuseDuplicateUsersPass(), InsertRescalePass(), + 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..39ca9151ee7 100644 --- a/backends/arm/test/passes/test_fuse_duplicate_users_pass.py +++ b/backends/arm/test/passes/test_fuse_duplicate_users_pass.py @@ -7,14 +7,23 @@ import executorch.backends.arm.tosa.dialect # noqa: F401 import torch -from executorch.backends.arm._passes import FuseDuplicateUsersPass +from executorch.backends.arm._passes import ( + EnsureUniqueOutputNodesPass, + FuseDuplicateUsersPass, + InsertRescalePass, + RemoveNoopPass, +) +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 +176,56 @@ 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() + + pass_manager = ArmPassManager(TosaCompileSpec("TOSA-1.0+FP")) + graph_module = pass_manager.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) + + pass_types = [type(pass_) for pass_ in pass_manager.passes] + post_noop_index = max( + index + for index, pass_type in enumerate(pass_types) + if pass_type is RemoveNoopPass + ) + assert pass_types[post_noop_index + 1 : post_noop_index + 4] == [ + FuseDuplicateUsersPass, + InsertRescalePass, + EnsureUniqueOutputNodesPass, + ]