Skip to content
Merged
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
3 changes: 3 additions & 0 deletions backends/arm/_passes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,9 @@
from .match_arg_dtype_pass import MatchArgDtypePass # noqa
from .match_arg_ranks_pass import MatchArgRanksPass # noqa
from .mm_to_bmm_pass import ConvertMmToBmmPass # noqa
from .move_data_movement_ops_to_smaller_dtype_pass import ( # noqa
MoveDataMovementOpsToSmallerDtypePass,
)
from .normalize_delegate_io_layout_pass import NormalizeDelegateIOLayoutPass # noqa
from .normalize_index_put_bool_index_tensor_pass import ( # noqa
NormalizeIndexPutBoolIndexTensorPass,
Expand Down
2 changes: 2 additions & 0 deletions backends/arm/_passes/arm_pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@
InsertTableOpsPass,
MatchArgDtypePass,
MatchArgRanksPass,
MoveDataMovementOpsToSmallerDtypePass,
NormalizeDelegateIOLayoutPass,
NormalizeIndexPutBoolIndexTensorPass,
NormalizeIndexPutNoneIndicesPass,
Expand Down Expand Up @@ -643,6 +644,7 @@ def _tosa_pipeline(
PropagateViewCopyPermuteUpPass(self.compile_spec, exported_program),
# Propagation can leave a binary op with mismatched operand ranks,
# which TOSA rejects; re-match ranks before lowering.
MoveDataMovementOpsToSmallerDtypePass(),
MatchArgRanksPass(exported_program),
RewriteHighRankSingletonPermutePass(),
DecomposePermuteForU55Pass(),
Expand Down
2 changes: 1 addition & 1 deletion backends/arm/_passes/decompose_var_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ def call_operator(self, op, args, kwargs, meta):
shape = [1 for _ in input_shape]

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

if op == torch.ops.aten.var.dim:
keepdim = False
Expand Down
71 changes: 71 additions & 0 deletions backends/arm/_passes/dim_maps.py
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,53 @@ def map_dim_inverse(
return None
return source_dims

def map_reduction_after_view(
self,
source_shape: Sequence[_Dim],
source_dims: int | Sequence[int],
) -> tuple[list[_Dim], list[int]] | None:
"""Map ``reduce(view(x), dims)`` to ``view(reduce(x, mapped_dims))``.

Returns the new view shape and reduction dims for:

view(reduce(x, source_dims), self.target_shape)
== reduce(view(x, new_shape), target_dims)

"""
target_shape = self.remap_target_shape(source_shape)
if target_shape is None:
return None

target_dims = self.map_dim(source_dims)
if target_dims is None or not self._is_contiguous_nonempty(target_dims):
return None
return target_shape, target_dims

def map_reduction_before_view(
self,
target_dims: int | Sequence[int],
) -> tuple[list[int], list[_Dim]] | None:
"""Map ``view(reduce(x, dims))`` to ``reduce(view(x), mapped_dims)``.

Returns the reduction dims and output view shape for:

reduce(view(x, self.target_shape), target_dims)
== view(reduce(x, source_dims), output_shape)

"""
source_dims = self.map_dim_inverse(target_dims)
if source_dims is None or not self._is_contiguous_nonempty(source_dims):
return None

try:
normalized_target_dims = _normalize_dims(target_dims, self.target_rank)
except AssertionError:
return None

return source_dims, self._reduce_shape(
self.target_shape, normalized_target_dims
)

def map_permutation(
self,
source_permutation: Sequence[int],
Expand Down Expand Up @@ -446,6 +493,8 @@ def map_permutation_inverse(
)

def remap_target_shape(self, source_shape: Sequence[_Dim]) -> list[_Dim] | None:
if not self.is_valid_map:
return None
if len(source_shape) != self.source_rank:
return None

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

if not same_numel(source_shape, target_shape):
return None
if self._has_zero_dim(target_shape):
return None
if not self._preserves_source_axis_order(source_shape, source_to_target_axes):
return None
return target_shape
Expand Down Expand Up @@ -551,6 +602,8 @@ def remap_unit_slice(
for target_axes in source_to_target_axes[:slice_dim]
for target_axis in target_axes
]
if not prev_target_axes:
return None
next_target_axes = [
target_axis
for target_axes in source_to_target_axes[slice_dim + 1 :]
Expand Down Expand Up @@ -810,6 +863,24 @@ def _is_valid_reduction_or_singleton(
group_to_axes[group].issubset(normalized_dims) for group in selected_groups
)

@staticmethod
def _is_contiguous_nonempty(dims: Sequence[int]) -> bool:
sorted_dims = sorted(set(dims))
return bool(sorted_dims) and sorted_dims == list(
range(sorted_dims[0], sorted_dims[-1] + 1)
)

@staticmethod
def _reduce_shape(shape: Sequence[_Dim], dims: Sequence[int]) -> list[_Dim]:
reduced_shape = list(shape)
for dim in dims:
reduced_shape[dim] = 1
return reduced_shape

@staticmethod
def _has_zero_dim(shape: Sequence[_Dim]) -> bool:
return any(_dim_equals(dim, 0) for dim in shape)

@classmethod
def _build_groups(
cls, source_shape: Sequence[_Dim], target_shape: Sequence[_Dim]
Expand Down
124 changes: 82 additions & 42 deletions backends/arm/_passes/fuse_identical_input_transforms_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,12 @@
import torch
from executorch.backends.arm._passes.arm_pass import ArmOpTargetedPass
from executorch.backends.arm._passes.arm_pass_utils import refresh_permute_view_meta
from executorch.backends.arm._passes.dim_maps import PermuteMap, same_numel, ViewMap
from executorch.backends.arm._passes.dim_maps import (
_dim_equals,
PermuteMap,
same_numel,
ViewMap,
)
from executorch.exir.dialects._ops import ops as exir_ops
from executorch.exir.pass_base import ExportPass, PassResult
from torch.export.exported_program import ExportedProgram
Expand Down Expand Up @@ -142,8 +147,12 @@ class FuseIdenticalInputTransformsPass(ArmOpTargetedPass):
exir_ops.edge.aten.bitwise_xor.Tensor,
exir_ops.edge.aten.remainder.Tensor,
}
_NARY_ELEMENTWISE_OPS = {
exir_ops.edge.aten.where.self,
}
_ELEMENTWISE_OPS = _BINARY_ELEMENTWISE_OPS | _NARY_ELEMENTWISE_OPS

target_ops = _BINARY_ELEMENTWISE_OPS | _CONCAT_OPS
target_ops = _ELEMENTWISE_OPS | _CONCAT_OPS

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

input_transforms = self._input_transforms(node)
if input_transforms is None:
input_nodes = list(node.all_input_nodes)
if len(input_nodes) < 2:
return False

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

transform = input_transforms[0]
updated_args = self._updated_node_args(
node, transform, node_val, input_transforms
)
transforms = [n for n in input_nodes if n.target in self._TARGETS]
if not transforms:
return False
transform = transforms[0]
if not self._inputs_share_transform_or_are_layout_invariant(
node, transform, input_nodes
):
return False

updated_args = self._updated_node_args(node, transform, node_val, input_nodes)
if updated_args is None:
return False
node_args, node_kwargs, transform_args, node_output_shape = updated_args

# Remove input transforms
producers = [n.all_input_nodes[0] for n in input_transforms]
producers = [
n.all_input_nodes[0] if n.target in self._TARGETS else n
for n in input_nodes
]

node.args = node_args
node.kwargs = node_kwargs
for input_transform, producer in zip(input_transforms, producers):
for input_transform, producer in zip(input_nodes, producers):
node.replace_input_with(input_transform, producer)
for input_transform in dict.fromkeys(input_transforms):
if len(input_transform.users) == 0:
for input_transform in dict.fromkeys(input_nodes):
if (
input_transform.target in self._TARGETS
and len(input_transform.users) == 0
):
node.graph.erase_node(input_transform)

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

def _updated_node_args(
self, node: Node, transform: Node, node_val: Any, input_transforms: list[Node]
self, node: Node, transform: Node, node_val: Any, input_nodes: list[Node]
) -> (
tuple[tuple[Any, ...], dict[str, Any], tuple[Any, ...], tuple[Any, ...]] | None
):
if not self._transforms_are_identical(input_transforms):
return None
if not self._transforms_only_used_by_node(node, input_transforms):
return None

if node.target in self._BINARY_ELEMENTWISE_OPS:
return self._update_node_args_binary(
node, transform, node_val, input_transforms
)
return self._update_node_args_binary(node, transform, node_val, input_nodes)

if node.target in self._CONCAT_OPS:
return self._update_node_args_concat(
node, transform, node_val, input_transforms
)
return self._update_node_args_concat(node, transform, node_val, input_nodes)

if node.target in self._NARY_ELEMENTWISE_OPS:
return self._update_node_args_binary(node, transform, node_val, input_nodes)

return None

def _inputs_share_transform_or_are_layout_invariant(
self, node: Node, transform: Node, input_nodes: list[Node]
) -> bool:
transforms = [n for n in input_nodes if n.target in self._TARGETS]
if not self._transforms_are_identical(transforms):
return False
if not self._transforms_only_used_by_node(node, transforms):
return False
if len(transforms) == len(input_nodes):
return True
if node.target not in self._ELEMENTWISE_OPS:
return False

transform_val = transform.meta.get("val")
if not isinstance(transform_val, torch.Tensor):
return False
rank = len(transform_val.shape)
return all(
input_node in transforms or self.is_layout_invariant(input_node, rank)
for input_node in input_nodes
)

@staticmethod
def is_layout_invariant(node: Node, rank: int) -> bool:
value = node.meta.get("val")
return (
isinstance(value, torch.Tensor)
and len(value.shape) == rank
and all(_dim_equals(dim, 1) for dim in value.shape)
)

def _transforms_are_identical(self, input_transforms: list[Node]) -> bool:
target = input_transforms[0].target
if target not in self._TARGETS:
Expand All @@ -281,7 +327,15 @@ def _transforms_only_used_by_node(

def _update_node_args_binary(self, node, transform, node_val, input_transforms):
producer_shapes = [
tuple(input_node.all_input_nodes[0].meta["val"].shape)
tuple(
(
input_node.all_input_nodes[0]
if input_node.target in self._TARGETS
else input_node
)
.meta["val"]
.shape
)
for input_node in input_transforms
]

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

Expand Down Expand Up @@ -382,20 +439,3 @@ def _mapped_concat_dim(self, transform: Node, concat_dim: int) -> int | None:
if mapped_dims is None or len(mapped_dims) != 1:
return None
return mapped_dims[0]

def _input_transforms(self, node: Node) -> list[Node] | None:
if node.target in self._BINARY_ELEMENTWISE_OPS:
input_transforms = list(node.args[:2])
elif node.target in self._CONCAT_OPS:
if len(node.args) == 0 or not isinstance(node.args[0], Sequence):
return None
input_transforms = list(node.args[0])
else:
return None

if len(input_transforms) < 2 or not all(
isinstance(n, Node) for n in input_transforms
):
return None

return cast(list[Node], input_transforms)
17 changes: 17 additions & 0 deletions backends/arm/_passes/match_arg_ranks_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,23 @@ def __init__(self, exported_program: ExportedProgram, *args, **kwargs) -> None:
exir_ops.edge.aten.bitwise_or.Tensor,
exir_ops.edge.aten.maximum.default,
exir_ops.edge.aten.minimum.default,
exir_ops.backend.tosa.ADD.default,
exir_ops.backend.tosa.ARITHMETIC_RIGHT_SHIFT.default,
exir_ops.backend.tosa.BITWISE_AND.default,
exir_ops.backend.tosa.BITWISE_OR.default,
exir_ops.backend.tosa.BITWISE_XOR.default,
exir_ops.backend.tosa.EQUAL.default,
exir_ops.backend.tosa.GREATER.default,
exir_ops.backend.tosa.GREATER_EQUAL.default,
exir_ops.backend.tosa.LOGICAL_AND.default,
exir_ops.backend.tosa.LOGICAL_LEFT_SHIFT.default,
exir_ops.backend.tosa.LOGICAL_OR.default,
exir_ops.backend.tosa.LOGICAL_XOR.default,
exir_ops.backend.tosa.MAXIMUM.default,
exir_ops.backend.tosa.MINIMUM.default,
exir_ops.backend.tosa.MUL.default,
exir_ops.backend.tosa.POW.default,
exir_ops.backend.tosa.SUB.default,
]

def _match_op_rank(self, graph_module, node, arg, max_rank):
Expand Down
Loading
Loading