Skip to content

Commit fb5b788

Browse files
Reland "Arm backend: Rerun duplicate-user fusion after TOSA lowering" (pytorch#21380)
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 <yufeng.shi@arm.com>
1 parent 3aeaa80 commit fb5b788

2 files changed

Lines changed: 67 additions & 2 deletions

File tree

backends/arm/_passes/arm_pass_manager.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -667,9 +667,12 @@ def _tosa_pipeline(
667667
SymbolicToTosaShapesPass(),
668668
InsertDynamicPaddingPass(),
669669
FuseConsecutiveConcatShapesPass(),
670-
EnsureUniqueOutputNodesPass(),
671670
RemoveNoopPass(),
671+
# Fuse duplicates exposed by late rewrites before inserting rescales;
672+
# fusing generated RESCALE users can corrupt distinct quantized paths.
673+
FuseDuplicateUsersPass(),
672674
InsertRescalePass(),
675+
EnsureUniqueOutputNodesPass(),
673676
]
674677
)
675678

backends/arm/test/passes/test_fuse_duplicate_users_pass.py

Lines changed: 63 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,14 +7,23 @@
77

88
import executorch.backends.arm.tosa.dialect # noqa: F401
99
import torch
10-
from executorch.backends.arm._passes import FuseDuplicateUsersPass
10+
from executorch.backends.arm._passes import (
11+
EnsureUniqueOutputNodesPass,
12+
FuseDuplicateUsersPass,
13+
InsertRescalePass,
14+
RemoveNoopPass,
15+
)
16+
from executorch.backends.arm._passes.arm_pass_manager import ArmPassManager
1117
from executorch.backends.arm.test import common
1218
from executorch.backends.arm.test.tester.test_pipeline import PassPipeline
19+
from executorch.backends.arm.tosa.compile_spec import TosaCompileSpec
1320
from executorch.backends.arm.tosa.specification import (
1421
TosaLoweringContext,
1522
TosaSpecification,
1623
)
24+
from executorch.exir import EdgeCompileConfig, to_edge
1725
from executorch.exir.dialects._ops import ops as exir_ops
26+
from torch.export import export
1827
from torch.fx import Graph, GraphModule
1928

2029
input_t = Tuple[torch.Tensor] # Input x
@@ -167,3 +176,56 @@ def test_fuse_duplicate_users_removes_identical_rescale_users():
167176
assert len(rescale_nodes) == 1
168177
output_node = result.graph_module.graph.output_node()
169178
assert output_node.args[0] == (rescale_nodes[0], rescale_nodes[0])
179+
180+
181+
class LateDuplicateUsers(torch.nn.Module):
182+
def __init__(self):
183+
super().__init__()
184+
self.register_buffer("first", torch.ones(2, 3))
185+
self.register_buffer("second", torch.ones(2, 3))
186+
187+
def forward(self, x):
188+
return x + self.first, x + self.second
189+
190+
191+
def test_fuse_duplicate_users_runs_after_tosa_transformations():
192+
exported_program = export(LateDuplicateUsers(), (torch.ones(2, 3),), strict=True)
193+
edge_program = to_edge(
194+
exported_program,
195+
compile_config=EdgeCompileConfig(_check_ir_validity=False),
196+
)
197+
edge_exported_program = edge_program.exported_program()
198+
199+
pass_manager = ArmPassManager(TosaCompileSpec("TOSA-1.0+FP"))
200+
graph_module = pass_manager.transform_to_backend_pipeline(
201+
edge_exported_program, edge_exported_program.graph_module
202+
)
203+
204+
add_nodes = [
205+
node
206+
for node in graph_module.graph.nodes
207+
if node.target == exir_ops.backend.tosa.ADD.default
208+
]
209+
identity_nodes = [
210+
node
211+
for node in graph_module.graph.nodes
212+
if node.target == exir_ops.backend.tosa.IDENTITY.default
213+
]
214+
215+
graph_module.graph.lint()
216+
assert len(add_nodes) == 1
217+
assert len(identity_nodes) == 2
218+
assert all(node.args[0] is add_nodes[0] for node in identity_nodes)
219+
assert graph_module.graph.output_node().args[0] == tuple(identity_nodes)
220+
221+
pass_types = [type(pass_) for pass_ in pass_manager.passes]
222+
post_noop_index = max(
223+
index
224+
for index, pass_type in enumerate(pass_types)
225+
if pass_type is RemoveNoopPass
226+
)
227+
assert pass_types[post_noop_index + 1 : post_noop_index + 4] == [
228+
FuseDuplicateUsersPass,
229+
InsertRescalePass,
230+
EnsureUniqueOutputNodesPass,
231+
]

0 commit comments

Comments
 (0)