Skip to content

Commit 8e653a6

Browse files
Arm backend: Cleanup dim-order and permute handling (pytorch#19278)
- Replace u55 permute dimension check with a u55-only pass decomposing large permutes. This pass checks for support by compiling targeted permutes using Vela to ensure alignment between Executorch and Vela. - Remove passes and testing not required anymore after dim-order update. - Remove all outdated mention of dim-order in the arm backend. Signed-off-by: Adrian Lundell <adrian.lundell@arm.com>
1 parent 0a113f8 commit 8e653a6

33 files changed

Lines changed: 439 additions & 1994 deletions

backends/arm/_passes/__init__.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,6 @@
77
from . import arm_pass_utils # noqa
88
from .arm_pass import ArmPass # noqa # usort: skip
99
from .accumulate_index_put_pass import AccumulateIndexPutPass # noqa
10-
from .annotate_output_dim_order_pass import AnnotateOutputDimOrderPass # noqa
1110
from .broadcast_args_pass import BroadcastArgsPass # noqa
1211
from .canonicalize_gather_pass import CanonicalizeGatherPass # noqa
1312
from .cast_int64_pass import CastInt64BuffersToInt32Pass # noqa
@@ -61,9 +60,6 @@
6160
from .decompose_index_tensor_to_gather_pass import ( # noqa
6261
DecomposeIndexTensorToGatherPass,
6362
)
64-
from .decompose_int16_activation_conv_pass import ( # noqa
65-
DecomposeConvWithInt16ActivationPass,
66-
)
6763
from .decompose_int_pow_pass import DecomposeIntPowPass # noqa
6864
from .decompose_layernorm_pass import DecomposeLayerNormPass # noqa
6965
from .decompose_leaky_relu_pass import DecomposeLeakyReLUPass # noqa
@@ -77,6 +73,7 @@
7773
from .decompose_maxpool2d_with_dilation_pass import DecomposeMaxPool2dPass # noqa
7874
from .decompose_meandim_pass import DecomposeMeanDimPass # noqa
7975
from .decompose_ne_pass import DecomposeNotEqualPass # noqa
76+
from .decompose_permute_for_u55_pass import DecomposePermuteForU55Pass # noqa
8077
from .decompose_quant_nodes import DecomposeQuantNodesPass # noqa
8178
from .decompose_remainder_pass import DecomposeRemainderPass # noqa
8279
from .decompose_rnn_pass import DecomposeRnnPass # noqa
@@ -167,7 +164,6 @@
167164
from .rewrite_upsample import RewriteUpsamplePass # noqa
168165
from .scalars_to_attribute_pass import ScalarsToAttributePass # noqa
169166
from .size_adjust_input_pass import SizeAdjustInputPass # noqa
170-
from .to_tosa_memory_format_pass import ToTosaMemoryFormatPass # noqa
171167
from .unsqueeze_before_repeat_pass import UnsqueezeBeforeRepeatPass # noqa
172168
from .unsqueeze_scalar_placeholders_pass import UnsqueezeScalarPlaceholdersPass # noqa
173169
from .replace_inf_and_limit_values_pass import ( # noqa # usort: skip

backends/arm/_passes/annotate_output_dim_order_pass.py

Lines changed: 0 additions & 28 deletions
This file was deleted.

backends/arm/_passes/arm_pass_manager.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,7 @@
7575
DecomposeMaxPool2dPass,
7676
DecomposeMeanDimPass,
7777
DecomposeNotEqualPass,
78+
DecomposePermuteForU55Pass,
7879
DecomposeQuantNodesPass,
7980
DecomposeRemainderPass,
8081
DecomposeRnnPass,
@@ -536,13 +537,14 @@ def _tosa_pipeline(
536537
RewriteConvPass(exported_program),
537538
RewriteMatmulPass(),
538539
RewritePadPass(),
539-
RewriteSlicePass(),
540540
FuseViewCopyTransformPass(),
541541
RemovePermutesAroundElementwiseOps(),
542542
PostponePermuteOpBelowSqueezeOrUnsqueezeLikeView(),
543543
FuseCascadedTransposeOrPermuteOps(),
544544
ConvertPermuteSingletonToViewPass(),
545545
RewriteHighRankSingletonPermutePass(),
546+
DecomposePermuteForU55Pass(),
547+
RewriteSlicePass(),
546548
InsertConstShapesPass(),
547549
]
548550
)

backends/arm/_passes/arm_pass_utils.py

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -364,11 +364,6 @@ def set_node_arg(node: torch.fx.Node, i: int | str, value):
364364
raise RuntimeError("Invalid type")
365365

366366

367-
def get_output_dim_orders(graph_module):
368-
output_node = graph_module.graph.output_node()
369-
return [get_first_fake_tensor(node).dim_order() for node in output_node.args[0]]
370-
371-
372367
def is_nested_control_flow_graph(graph_module: GraphModule) -> bool:
373368
"""Returns True if graph_module is a nested control-flow graph."""
374369

backends/arm/_passes/decompose_int16_activation_conv_pass.py

Lines changed: 0 additions & 147 deletions
This file was deleted.

0 commit comments

Comments
 (0)