Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
105 changes: 103 additions & 2 deletions backends/arm/_passes/rewrite_conv_pass.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
# Copyright 2025-2026 Arm Limited and/or its affiliates.
#
# This source code is licensed under the BSD-style license found in the
Expand Down Expand Up @@ -37,11 +37,20 @@
TOSA_CONTROL_FLOW_SOURCE_NODE_META,
TosaSpecialDtype,
)
from executorch.backends.arm.tosa.specification import get_context_shape_env
from executorch.backends.arm.tosa.specification import (
get_context_shape_env,
get_context_spec,
)
from executorch.backends.transforms.fuse_duplicate_users_pass import (
build_node_signature,
DO_NOT_FUSE_DUPLICATE_META_KEY,
)
from executorch.backends.transforms.utils import create_constant_placeholder
from executorch.exir.dialects._ops import ops as exir_ops
from executorch.exir.dialects.edge._ops import EdgeOpOverload
from executorch.exir.pass_base import ExportPass, PassResult

from torch._ops import OpOverload
from torch._subclasses.fake_tensor import FakeTensor
from torch.export.graph_signature import InputKind

Expand Down Expand Up @@ -533,7 +542,7 @@

def _get_direct_int32_rescale_users(
self, node: torch.fx.Node
) -> list[torch.fx.Node]:
) -> list[torch.fx.Node] | None:
"""Return consumers that directly request an INT32 value."""
return [user for user in node.users if self._is_direct_int32_rescale(user)]

Expand Down Expand Up @@ -571,6 +580,87 @@
output.meta["val"] = output_fake_tensor
return output, output_fake_tensor

@classmethod
def _deduplicate_a16w8_output_rescales(
cls,
graph_module: torch.fx.GraphModule,
tosa_op: torch.fx.Node,
node_order: dict[torch.fx.Node, int],
) -> list[torch.fx.Node]:
"""Merge only complete, canonical RESCALE-to-PERMUTE heads."""
rescale_users = sorted(
tosa_op.users, key=lambda node: node_order.get(node, len(node_order))
)
if any(
user.target != exir_ops.backend.tosa.RESCALE.default
for user in rescale_users
):
return None

Check failure on line 598 in backends/arm/_passes/rewrite_conv_pass.py

View workflow job for this annotation

GitHub Actions / lintrunner-mypy

MYPY return-value

Incompatible return value type (got "None", expected "list[Node]") To disable, use ` # type: ignore[return-value]`

unique_rescales: dict[tuple[Any, ...], torch.fx.Node] = {}
deduplicated_rescales: list[torch.fx.Node] = []
for rescale in rescale_users:
rescale_outputs = list(rescale.users)
if (
len(rescale_outputs) != 1
or rescale_outputs[0].target != exir_ops.edge.aten.permute_copy.default
):
deduplicated_rescales.append(rescale)
continue
layout_permute = rescale_outputs[0]
rescale_signature = build_node_signature(
rescale, positional_arg_start=1
)
permute_signature = build_node_signature(
layout_permute, positional_arg_start=1
)
if rescale_signature is None or permute_signature is None:
deduplicated_rescales.append(rescale)
continue
signature = (
rescale_signature,
permute_signature,
)
canonical_permute = unique_rescales.get(signature)
if canonical_permute is not None:
# Layout permutes are inserted directly after their RESCALE,
# so the earliest RESCALE also provides a dominating permute.
layout_permute.replace_all_uses_with(canonical_permute)
graph_module.graph.erase_node(layout_permute)
graph_module.graph.erase_node(rescale)
else:
unique_rescales[signature] = layout_permute
deduplicated_rescales.append(rescale)

return deduplicated_rescales

def _separate_u55_a16w8_output_rescales(
self,
graph_module: torch.fx.GraphModule,
tosa_op: torch.fx.Node,
node_order: dict[torch.fx.Node, int],
) -> None:
if len(tosa_op.users) < 2:
return

rescale_users = self._deduplicate_a16w8_output_rescales(
graph_module, tosa_op, node_order
)
if rescale_users is None or len(rescale_users) < 2:
return
tosa_op.meta[DO_NOT_FUSE_DUPLICATE_META_KEY] = True
for rescale in rescale_users[1:]:
with graph_module.graph.inserting_before(rescale):
cloned_tosa_op = create_node(
graph=graph_module.graph,
op_target=cast(OpOverload | EdgeOpOverload, tosa_op.target),
args=tosa_op.args,
kwargs=tosa_op.kwargs,
from_node=tosa_op,
inherit_qparams=True,
)
rescale.replace_input_with(tosa_op, cloned_tosa_op)

def _insert_a16w8_output_branches(
self,
graph_module: torch.fx.GraphModule,
Expand Down Expand Up @@ -611,7 +701,7 @@
#
# Direct consumers have no axis-changing operation between the
# convolution and rescale, so both scale factors can be combined.
for int32_user in direct_int32_users:

Check failure on line 704 in backends/arm/_passes/rewrite_conv_pass.py

View workflow job for this annotation

GitHub Actions / lintrunner-mypy

MYPY union-attr

Item "None" of "list[Node] | None" has no attribute "__iter__" (not iterable) To disable, use ` # type: ignore[union-attr]`
# Retarget the existing consumer rescale to the TOSA convolution;
# its previous consumers will be moved after the layout permute.
user_scales = cast(list[float], int32_user.args[2])
Expand Down Expand Up @@ -787,6 +877,7 @@

def call(self, graph_module: torch.fx.GraphModule) -> PassResult: # noqa: C901
modified = False
a16w8_tosa_ops: list[torch.fx.Node] = []
for node in graph_module.graph.nodes:
if (
node.op != "call_function"
Expand Down Expand Up @@ -1172,11 +1263,21 @@
if squeeze_view is not None:
graph_module.graph.erase_node(squeeze_view)
graph_module.graph.erase_node(output_conversion_node)
a16w8_tosa_ops.append(tosa_op)
else:
node.replace_all_uses_with(node_replacement)

graph_module.graph.erase_node(node)

if a16w8_tosa_ops and get_context_spec().is_U55_subset:
node_order = {
node: index for index, node in enumerate(graph_module.graph.nodes)
}
for tosa_op in a16w8_tosa_ops:
self._separate_u55_a16w8_output_rescales(
graph_module, tosa_op, node_order
)

if modified:
graph_module.recompile()
graph_module = super().call(graph_module).graph_module
Expand Down
19 changes: 19 additions & 0 deletions backends/arm/test/passes/test_fuse_duplicate_users_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,9 @@
TosaLoweringContext,
TosaSpecification,
)
from executorch.backends.transforms.fuse_duplicate_users_pass import (
DO_NOT_FUSE_DUPLICATE_META_KEY,
)
from executorch.exir import EdgeCompileConfig, to_edge
from executorch.exir.dialects._ops import ops as exir_ops
from torch.export import export
Expand Down Expand Up @@ -163,6 +166,22 @@ def test_fuse_duplicate_users_preserves_graph_order_for_representative():
assert len(_add_node_names(result.graph_module)) == 1


def test_fuse_duplicate_users_honors_do_not_fuse_marker():
graph_module = _graph_with_users_not_in_node_order()
marked_node = next(
node
for node in graph_module.graph.nodes
if node.target == torch.ops.aten.add.Tensor
)
marked_node.meta[DO_NOT_FUSE_DUPLICATE_META_KEY] = True

result = FuseDuplicateUsersPass()(graph_module)

result.graph_module.graph.lint()
assert not result.modified
assert len(_add_node_names(result.graph_module)) == 2


def test_fuse_duplicate_users_keeps_identical_rescale_users():
graph_module = _graph_with_duplicate_rescale_users()

Expand Down
Loading
Loading