Skip to content

Commit 179c4ee

Browse files
Arm backend: Improve permute/view fusing (#21055)
- Makes view/permutes always split to all branches of a node with multiple outputs for downward propagation. Note that this will never increase the number of permutes/views in the graph, since any non-fused permutes/views will propagate back and be fused together again by the upward propagation pass. - Adds support for where operator to fuse_identical_input_transforms - Replaces _would_strand_layout_op_on_wider_elements with MoveDataMovementOpsToSmallerDtypePass to allow ops to propagate more freely during the fusing phase - Adds tosa ops to MatchArgsPass since it is run after some ops have been rewritten to TOSA. - Generalizes refresh_permute_view_meta into refresh_node_meta which works for any operators and uses that. Signed-off-by: Adrian Lundell <adrian.lundell@arm.com>
1 parent 007fc2b commit 179c4ee

15 files changed

Lines changed: 1202 additions & 212 deletions

backends/arm/_passes/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -147,6 +147,9 @@
147147
from .match_arg_dtype_pass import MatchArgDtypePass # noqa
148148
from .match_arg_ranks_pass import MatchArgRanksPass # noqa
149149
from .mm_to_bmm_pass import ConvertMmToBmmPass # noqa
150+
from .move_data_movement_ops_to_smaller_dtype_pass import ( # noqa
151+
MoveDataMovementOpsToSmallerDtypePass,
152+
)
150153
from .normalize_delegate_io_layout_pass import NormalizeDelegateIOLayoutPass # noqa
151154
from .normalize_index_put_bool_index_tensor_pass import ( # noqa
152155
NormalizeIndexPutBoolIndexTensorPass,

backends/arm/_passes/arm_pass_manager.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,7 @@
129129
InsertTableOpsPass,
130130
MatchArgDtypePass,
131131
MatchArgRanksPass,
132+
MoveDataMovementOpsToSmallerDtypePass,
132133
NormalizeDelegateIOLayoutPass,
133134
NormalizeIndexPutBoolIndexTensorPass,
134135
NormalizeIndexPutNoneIndicesPass,
@@ -643,6 +644,7 @@ def _tosa_pipeline(
643644
PropagateViewCopyPermuteUpPass(self.compile_spec, exported_program),
644645
# Propagation can leave a binary op with mismatched operand ranks,
645646
# which TOSA rejects; re-match ranks before lowering.
647+
MoveDataMovementOpsToSmallerDtypePass(),
646648
MatchArgRanksPass(exported_program),
647649
RewriteHighRankSingletonPermutePass(),
648650
DecomposePermuteForU55Pass(),

backends/arm/_passes/decompose_var_pass.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,7 @@ def call_operator(self, op, args, kwargs, meta):
7474
shape = [1 for _ in input_shape]
7575

7676
# Get dim from args based on argument type
77-
dim = get_node_arg(args, key=list, default_value=list(range(len(shape))))
77+
dim = get_node_arg(args, key=list, default_value=list(range(len(input_shape))))
7878

7979
if op == torch.ops.aten.var.dim:
8080
keepdim = False

backends/arm/_passes/dim_maps.py

Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -363,6 +363,53 @@ def map_dim_inverse(
363363
return None
364364
return source_dims
365365

366+
def map_reduction_after_view(
367+
self,
368+
source_shape: Sequence[_Dim],
369+
source_dims: int | Sequence[int],
370+
) -> tuple[list[_Dim], list[int]] | None:
371+
"""Map ``reduce(view(x), dims)`` to ``view(reduce(x, mapped_dims))``.
372+
373+
Returns the new view shape and reduction dims for:
374+
375+
view(reduce(x, source_dims), self.target_shape)
376+
== reduce(view(x, new_shape), target_dims)
377+
378+
"""
379+
target_shape = self.remap_target_shape(source_shape)
380+
if target_shape is None:
381+
return None
382+
383+
target_dims = self.map_dim(source_dims)
384+
if target_dims is None or not self._is_contiguous_nonempty(target_dims):
385+
return None
386+
return target_shape, target_dims
387+
388+
def map_reduction_before_view(
389+
self,
390+
target_dims: int | Sequence[int],
391+
) -> tuple[list[int], list[_Dim]] | None:
392+
"""Map ``view(reduce(x, dims))`` to ``reduce(view(x), mapped_dims)``.
393+
394+
Returns the reduction dims and output view shape for:
395+
396+
reduce(view(x, self.target_shape), target_dims)
397+
== view(reduce(x, source_dims), output_shape)
398+
399+
"""
400+
source_dims = self.map_dim_inverse(target_dims)
401+
if source_dims is None or not self._is_contiguous_nonempty(source_dims):
402+
return None
403+
404+
try:
405+
normalized_target_dims = _normalize_dims(target_dims, self.target_rank)
406+
except AssertionError:
407+
return None
408+
409+
return source_dims, self._reduce_shape(
410+
self.target_shape, normalized_target_dims
411+
)
412+
366413
def map_permutation(
367414
self,
368415
source_permutation: Sequence[int],
@@ -446,6 +493,8 @@ def map_permutation_inverse(
446493
)
447494

448495
def remap_target_shape(self, source_shape: Sequence[_Dim]) -> list[_Dim] | None:
496+
if not self.is_valid_map:
497+
return None
449498
if len(source_shape) != self.source_rank:
450499
return None
451500

@@ -470,6 +519,8 @@ def remap_target_shape(self, source_shape: Sequence[_Dim]) -> list[_Dim] | None:
470519

471520
if not same_numel(source_shape, target_shape):
472521
return None
522+
if self._has_zero_dim(target_shape):
523+
return None
473524
if not self._preserves_source_axis_order(source_shape, source_to_target_axes):
474525
return None
475526
return target_shape
@@ -551,6 +602,8 @@ def remap_unit_slice(
551602
for target_axes in source_to_target_axes[:slice_dim]
552603
for target_axis in target_axes
553604
]
605+
if not prev_target_axes:
606+
return None
554607
next_target_axes = [
555608
target_axis
556609
for target_axes in source_to_target_axes[slice_dim + 1 :]
@@ -810,6 +863,24 @@ def _is_valid_reduction_or_singleton(
810863
group_to_axes[group].issubset(normalized_dims) for group in selected_groups
811864
)
812865

866+
@staticmethod
867+
def _is_contiguous_nonempty(dims: Sequence[int]) -> bool:
868+
sorted_dims = sorted(set(dims))
869+
return bool(sorted_dims) and sorted_dims == list(
870+
range(sorted_dims[0], sorted_dims[-1] + 1)
871+
)
872+
873+
@staticmethod
874+
def _reduce_shape(shape: Sequence[_Dim], dims: Sequence[int]) -> list[_Dim]:
875+
reduced_shape = list(shape)
876+
for dim in dims:
877+
reduced_shape[dim] = 1
878+
return reduced_shape
879+
880+
@staticmethod
881+
def _has_zero_dim(shape: Sequence[_Dim]) -> bool:
882+
return any(_dim_equals(dim, 0) for dim in shape)
883+
813884
@classmethod
814885
def _build_groups(
815886
cls, source_shape: Sequence[_Dim], target_shape: Sequence[_Dim]

backends/arm/_passes/fuse_identical_input_transforms_pass.py

Lines changed: 82 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,12 @@
1111
import torch
1212
from executorch.backends.arm._passes.arm_pass import ArmOpTargetedPass
1313
from executorch.backends.arm._passes.arm_pass_utils import refresh_permute_view_meta
14-
from executorch.backends.arm._passes.dim_maps import PermuteMap, same_numel, ViewMap
14+
from executorch.backends.arm._passes.dim_maps import (
15+
_dim_equals,
16+
PermuteMap,
17+
same_numel,
18+
ViewMap,
19+
)
1520
from executorch.exir.dialects._ops import ops as exir_ops
1621
from executorch.exir.pass_base import ExportPass, PassResult
1722
from torch.export.exported_program import ExportedProgram
@@ -142,8 +147,12 @@ class FuseIdenticalInputTransformsPass(ArmOpTargetedPass):
142147
exir_ops.edge.aten.bitwise_xor.Tensor,
143148
exir_ops.edge.aten.remainder.Tensor,
144149
}
150+
_NARY_ELEMENTWISE_OPS = {
151+
exir_ops.edge.aten.where.self,
152+
}
153+
_ELEMENTWISE_OPS = _BINARY_ELEMENTWISE_OPS | _NARY_ELEMENTWISE_OPS
145154

146-
target_ops = _BINARY_ELEMENTWISE_OPS | _CONCAT_OPS
155+
target_ops = _ELEMENTWISE_OPS | _CONCAT_OPS
147156

148157
def __init__(self, exported_program: ExportedProgram | None = None) -> None:
149158
super().__init__()
@@ -180,31 +189,43 @@ def _sink_identical_input_transforms(self, node: Node) -> bool:
180189
if node.target not in self.target_ops:
181190
return False
182191

183-
input_transforms = self._input_transforms(node)
184-
if input_transforms is None:
192+
input_nodes = list(node.all_input_nodes)
193+
if len(input_nodes) < 2:
185194
return False
186195

187196
node_val = node.meta.get("val", None)
188197
if node_val is None:
189198
return False
190199

191-
transform = input_transforms[0]
192-
updated_args = self._updated_node_args(
193-
node, transform, node_val, input_transforms
194-
)
200+
transforms = [n for n in input_nodes if n.target in self._TARGETS]
201+
if not transforms:
202+
return False
203+
transform = transforms[0]
204+
if not self._inputs_share_transform_or_are_layout_invariant(
205+
node, transform, input_nodes
206+
):
207+
return False
208+
209+
updated_args = self._updated_node_args(node, transform, node_val, input_nodes)
195210
if updated_args is None:
196211
return False
197212
node_args, node_kwargs, transform_args, node_output_shape = updated_args
198213

199214
# Remove input transforms
200-
producers = [n.all_input_nodes[0] for n in input_transforms]
215+
producers = [
216+
n.all_input_nodes[0] if n.target in self._TARGETS else n
217+
for n in input_nodes
218+
]
201219

202220
node.args = node_args
203221
node.kwargs = node_kwargs
204-
for input_transform, producer in zip(input_transforms, producers):
222+
for input_transform, producer in zip(input_nodes, producers):
205223
node.replace_input_with(input_transform, producer)
206-
for input_transform in dict.fromkeys(input_transforms):
207-
if len(input_transform.users) == 0:
224+
for input_transform in dict.fromkeys(input_nodes):
225+
if (
226+
input_transform.target in self._TARGETS
227+
and len(input_transform.users) == 0
228+
):
208229
node.graph.erase_node(input_transform)
209230

210231
node.meta = copy.copy(node.meta)
@@ -235,27 +256,52 @@ def _new_transform_meta(self, node: Node, transform: Node) -> dict[str, Any]:
235256
return meta
236257

237258
def _updated_node_args(
238-
self, node: Node, transform: Node, node_val: Any, input_transforms: list[Node]
259+
self, node: Node, transform: Node, node_val: Any, input_nodes: list[Node]
239260
) -> (
240261
tuple[tuple[Any, ...], dict[str, Any], tuple[Any, ...], tuple[Any, ...]] | None
241262
):
242-
if not self._transforms_are_identical(input_transforms):
243-
return None
244-
if not self._transforms_only_used_by_node(node, input_transforms):
245-
return None
246-
247263
if node.target in self._BINARY_ELEMENTWISE_OPS:
248-
return self._update_node_args_binary(
249-
node, transform, node_val, input_transforms
250-
)
264+
return self._update_node_args_binary(node, transform, node_val, input_nodes)
251265

252266
if node.target in self._CONCAT_OPS:
253-
return self._update_node_args_concat(
254-
node, transform, node_val, input_transforms
255-
)
267+
return self._update_node_args_concat(node, transform, node_val, input_nodes)
268+
269+
if node.target in self._NARY_ELEMENTWISE_OPS:
270+
return self._update_node_args_binary(node, transform, node_val, input_nodes)
256271

257272
return None
258273

274+
def _inputs_share_transform_or_are_layout_invariant(
275+
self, node: Node, transform: Node, input_nodes: list[Node]
276+
) -> bool:
277+
transforms = [n for n in input_nodes if n.target in self._TARGETS]
278+
if not self._transforms_are_identical(transforms):
279+
return False
280+
if not self._transforms_only_used_by_node(node, transforms):
281+
return False
282+
if len(transforms) == len(input_nodes):
283+
return True
284+
if node.target not in self._ELEMENTWISE_OPS:
285+
return False
286+
287+
transform_val = transform.meta.get("val")
288+
if not isinstance(transform_val, torch.Tensor):
289+
return False
290+
rank = len(transform_val.shape)
291+
return all(
292+
input_node in transforms or self.is_layout_invariant(input_node, rank)
293+
for input_node in input_nodes
294+
)
295+
296+
@staticmethod
297+
def is_layout_invariant(node: Node, rank: int) -> bool:
298+
value = node.meta.get("val")
299+
return (
300+
isinstance(value, torch.Tensor)
301+
and len(value.shape) == rank
302+
and all(_dim_equals(dim, 1) for dim in value.shape)
303+
)
304+
259305
def _transforms_are_identical(self, input_transforms: list[Node]) -> bool:
260306
target = input_transforms[0].target
261307
if target not in self._TARGETS:
@@ -281,7 +327,15 @@ def _transforms_only_used_by_node(
281327

282328
def _update_node_args_binary(self, node, transform, node_val, input_transforms):
283329
producer_shapes = [
284-
tuple(input_node.all_input_nodes[0].meta["val"].shape)
330+
tuple(
331+
(
332+
input_node.all_input_nodes[0]
333+
if input_node.target in self._TARGETS
334+
else input_node
335+
)
336+
.meta["val"]
337+
.shape
338+
)
285339
for input_node in input_transforms
286340
]
287341

@@ -293,6 +347,9 @@ def _update_node_args_binary(self, node, transform, node_val, input_transforms):
293347
transform_args = (node, *transform.args[1:])
294348
if transform.target == self._VIEW_TARGET:
295349
transform_args = (node, list(node_val.shape))
350+
# Reshaping before an elementwise op can change which dimensions
351+
# broadcast. Sinking is safe only when the broadcast in the source
352+
# layout already has exactly one of the producer shapes.
296353
if node_output_shape not in producer_shapes:
297354
return None
298355

@@ -382,20 +439,3 @@ def _mapped_concat_dim(self, transform: Node, concat_dim: int) -> int | None:
382439
if mapped_dims is None or len(mapped_dims) != 1:
383440
return None
384441
return mapped_dims[0]
385-
386-
def _input_transforms(self, node: Node) -> list[Node] | None:
387-
if node.target in self._BINARY_ELEMENTWISE_OPS:
388-
input_transforms = list(node.args[:2])
389-
elif node.target in self._CONCAT_OPS:
390-
if len(node.args) == 0 or not isinstance(node.args[0], Sequence):
391-
return None
392-
input_transforms = list(node.args[0])
393-
else:
394-
return None
395-
396-
if len(input_transforms) < 2 or not all(
397-
isinstance(n, Node) for n in input_transforms
398-
):
399-
return None
400-
401-
return cast(list[Node], input_transforms)

backends/arm/_passes/match_arg_ranks_pass.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,23 @@ def __init__(self, exported_program: ExportedProgram, *args, **kwargs) -> None:
6666
exir_ops.edge.aten.bitwise_or.Tensor,
6767
exir_ops.edge.aten.maximum.default,
6868
exir_ops.edge.aten.minimum.default,
69+
exir_ops.backend.tosa.ADD.default,
70+
exir_ops.backend.tosa.ARITHMETIC_RIGHT_SHIFT.default,
71+
exir_ops.backend.tosa.BITWISE_AND.default,
72+
exir_ops.backend.tosa.BITWISE_OR.default,
73+
exir_ops.backend.tosa.BITWISE_XOR.default,
74+
exir_ops.backend.tosa.EQUAL.default,
75+
exir_ops.backend.tosa.GREATER.default,
76+
exir_ops.backend.tosa.GREATER_EQUAL.default,
77+
exir_ops.backend.tosa.LOGICAL_AND.default,
78+
exir_ops.backend.tosa.LOGICAL_LEFT_SHIFT.default,
79+
exir_ops.backend.tosa.LOGICAL_OR.default,
80+
exir_ops.backend.tosa.LOGICAL_XOR.default,
81+
exir_ops.backend.tosa.MAXIMUM.default,
82+
exir_ops.backend.tosa.MINIMUM.default,
83+
exir_ops.backend.tosa.MUL.default,
84+
exir_ops.backend.tosa.POW.default,
85+
exir_ops.backend.tosa.SUB.default,
6986
]
7087

7188
def _match_op_rank(self, graph_module, node, arg, max_rank):

0 commit comments

Comments
 (0)