diff --git a/backends/arm/_passes/fuse_duplicate_users_pass.py b/backends/arm/_passes/fuse_duplicate_users_pass.py index 9bd21112569..746dade94bd 100644 --- a/backends/arm/_passes/fuse_duplicate_users_pass.py +++ b/backends/arm/_passes/fuse_duplicate_users_pass.py @@ -45,9 +45,17 @@ def call(self, graph_module: GraphModule) -> PassResult: node_order = {node: index for index, node in enumerate(graph.nodes)} producers: Deque[Node] = deque(node for node in graph.nodes) + queued_producers: Set[Node] = set(producers) + + def enqueue_producer(node: Node) -> None: + if node.graph is None or node in queued_producers: + return + producers.append(node) + queued_producers.add(node) while producers: producer = producers.popleft() + queued_producers.discard(producer) if producer.graph is None: # Node was deleted by a previous rewrite while still queued. @@ -84,8 +92,8 @@ def call(self, graph_module: GraphModule) -> PassResult: # Revisit the current producer and the surviving user so that # newly formed duplicate chains can be fused in later # iterations. - producers.append(producer) - producers.append(representative) + enqueue_producer(producer) + enqueue_producer(representative) if modified: graph_module.recompile() diff --git a/backends/arm/_passes/propagate_view_copy_permute_pass.py b/backends/arm/_passes/propagate_view_copy_permute_pass.py index 938c47f1ec8..2b2f08ca46d 100644 --- a/backends/arm/_passes/propagate_view_copy_permute_pass.py +++ b/backends/arm/_passes/propagate_view_copy_permute_pass.py @@ -106,20 +106,25 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult: continue if self._propagate(node): iteration_modified = True + graph_module = self._retrace(graph_module) break if iteration_modified: - graph_module = self._retrace(graph_module) - result = self.fuse_horizontal(graph_module) - graph_module = result.graph_module - iteration_modified |= result.modified - result = self.fuse_vertical(graph_module) - graph_module = result.graph_module - iteration_modified |= result.modified + modified = True + continue + + result = self.fuse_horizontal(graph_module) + graph_module = result.graph_module + iteration_modified |= result.modified + result = self.fuse_vertical(graph_module) + graph_module = result.graph_module + iteration_modified |= result.modified modified |= iteration_modified - if not iteration_modified: - break + if iteration_modified: + graph_module = self._retrace(graph_module) + continue + break if modified: graph_module = self._retrace(graph_module) diff --git a/backends/transforms/canonicalize_view_copy_permute_pass.py b/backends/transforms/canonicalize_view_copy_permute_pass.py index 56dbc14bb0a..baa88191393 100644 --- a/backends/transforms/canonicalize_view_copy_permute_pass.py +++ b/backends/transforms/canonicalize_view_copy_permute_pass.py @@ -5,6 +5,7 @@ from __future__ import annotations +from collections import deque from typing import cast, Sequence, Set, Type import torch @@ -90,20 +91,25 @@ def _collect_chains(self, graph_module: GraphModule) -> list[list[Node]]: """Returns a list of linear chains of view/permutes in the graph.""" chains: list[list[Node]] = [] - view_permute_nodes = [ + view_permute_nodes = deque( node for node in graph_module.graph.nodes if node.target in self._TARGETS - ] + ) + remaining = set(view_permute_nodes) while view_permute_nodes: - node = view_permute_nodes.pop(0) + node = view_permute_nodes.popleft() + if node not in remaining: + continue + remaining.remove(node) + chain = [node] current = node while len(current.users) == 1: user = next(iter(current.users)) - if user.target not in self._TARGETS: + if user.target not in self._TARGETS or user not in remaining: break - view_permute_nodes.remove(user) + remaining.remove(user) chain.append(user) current = user