Skip to content

Commit fe0d927

Browse files
DrJessopmeta-codesync[bot]
authored andcommitted
Remove or move permute after mean
Summary: If we have a permute -> unary chain -> mean, based on the reduction dims of the mean, we can either fully remove the preceding permute or move the permute after the mean. Case 1: Dims after permute are still in same order with respect to each other, we can fully get rid of the permute and just update the reduction dims of the mean. Case 2: Not case 1. In this case, it's better to move the permute after the mean, since the permute will operate on less data. Differential Revision: D102268214
1 parent c3f3d12 commit fe0d927

2 files changed

Lines changed: 317 additions & 0 deletions

File tree

backends/cadence/aot/remove_ops.py

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -387,6 +387,127 @@ def maybe_remove_or_replace(self, node: Node) -> bool:
387387
return False
388388

389389

390+
@register_cadence_pass(CadencePassAttribute(opt_level=1))
391+
class RemovePermuteBeforeMeanPass(RemoveOrReplacePassInterface):
392+
"""Remove or sink permute ops that precede mean reductions through unary chains.
393+
394+
When a permute feeds into a mean (possibly through unary ops like
395+
dequantize/quantize), two optimizations apply:
396+
397+
1. If non-reduced dims maintain their relative order and positions, the
398+
permute is fully removed and the mean's reduction dims are remapped.
399+
2. Otherwise, the permute is moved after the mean so it operates on
400+
smaller data.
401+
"""
402+
403+
_UNARY_TARGETS: frozenset[EdgeOpOverload] = frozenset(
404+
{
405+
exir_ops.edge.cadence.dequantize_per_tensor.default,
406+
exir_ops.edge.cadence.quantize_per_tensor.default,
407+
exir_ops.edge.quantized_decomposed.dequantize_per_tensor.default,
408+
exir_ops.edge.quantized_decomposed.quantize_per_tensor.default,
409+
exir_ops.edge.aten.clone.default,
410+
exir_ops.edge.aten.relu.default,
411+
exir_ops.edge.aten.neg.default,
412+
exir_ops.edge.aten.abs.default,
413+
}
414+
)
415+
416+
@property
417+
def targets(self) -> list[EdgeOpOverload]:
418+
return [exir_ops.edge.aten.mean.dim]
419+
420+
def _find_permute_through_unary_chain(self, mean_node: Node) -> Optional[Node]:
421+
"""Walk backward from mean through single-user unary ops to find a permute."""
422+
current = mean_node.args[0]
423+
if not isinstance(current, Node):
424+
return None
425+
while True:
426+
if current.target == exir_ops.edge.aten.permute_copy.default:
427+
return current
428+
if current.target not in self._UNARY_TARGETS:
429+
return None
430+
if len(current.users) != 1:
431+
return None
432+
parent = current.args[0]
433+
if not isinstance(parent, Node):
434+
return None
435+
current = parent
436+
437+
@staticmethod
438+
def _get_keepdim(node: Node) -> bool:
439+
if len(node.args) >= 3:
440+
return bool(node.args[2])
441+
return bool(node.kwargs.get("keepdim", False))
442+
443+
@staticmethod
444+
def _can_fully_remove(
445+
perm: list[int], new_reduction_dims: list[int], ndim: int, keepdim: bool
446+
) -> bool:
447+
"""Check whether the post-mean permute would be a no-op."""
448+
canonical_reduction = {d % ndim for d in new_reduction_dims}
449+
if keepdim:
450+
return all(
451+
perm[d] == d for d in range(ndim) if d not in canonical_reduction
452+
)
453+
non_reduced_in_perm_order = [d for d in perm if d not in canonical_reduction]
454+
return non_reduced_in_perm_order == sorted(non_reduced_in_perm_order)
455+
456+
@staticmethod
457+
def _compute_post_mean_perm(
458+
perm: list[int], new_reduction_dims: list[int], ndim: int, keepdim: bool
459+
) -> list[int]:
460+
"""Compute the permutation to insert after the mean."""
461+
if keepdim:
462+
return list(perm)
463+
canonical_reduction = {d % ndim for d in new_reduction_dims}
464+
non_reduced_original = sorted(
465+
d for d in range(ndim) if d not in canonical_reduction
466+
)
467+
non_reduced_permuted = [d for d in perm if d not in canonical_reduction]
468+
return [non_reduced_original.index(d) for d in non_reduced_permuted]
469+
470+
def maybe_remove_or_replace(self, node: Node) -> bool:
471+
reduction_dims = cast(list[int], node.args[1])
472+
473+
permute_node = self._find_permute_through_unary_chain(node)
474+
if permute_node is None:
475+
return False
476+
477+
perm = cast(list[int], permute_node.args[1])
478+
ndim = len(perm)
479+
480+
if len(permute_node.users) != 1:
481+
return False
482+
483+
permute_input = permute_node.args[0]
484+
assert isinstance(permute_input, Node)
485+
486+
new_reduction_dims = [perm[d % ndim] for d in reduction_dims]
487+
keepdim = self._get_keepdim(node)
488+
can_remove = self._can_fully_remove(perm, new_reduction_dims, ndim, keepdim)
489+
490+
permute_node.replace_all_uses_with(permute_input)
491+
node.args = (node.args[0], new_reduction_dims) + node.args[2:]
492+
493+
if not can_remove:
494+
post_perm = self._compute_post_mean_perm(
495+
perm, new_reduction_dims, ndim, keepdim
496+
)
497+
graph = node.graph
498+
with graph.inserting_after(node):
499+
new_permute = graph.create_node(
500+
"call_function",
501+
exir_ops.edge.aten.permute_copy.default,
502+
args=(node, post_perm),
503+
)
504+
for user in list(node.users):
505+
if user is not new_permute:
506+
user.replace_input_with(node, new_permute)
507+
508+
return True
509+
510+
390511
@register_cadence_pass(CadencePassAttribute(opt_level=2))
391512
class RemovePermutesAroundElementwiseOps(_SharedRemovePermutesAroundElementwiseOps):
392513
permutable_ops: set[EdgeOpOverload] = (
@@ -646,6 +767,7 @@ class CommonRemovePasses:
646767
RemoveNopSliceOrViewOpPass,
647768
RemoveToOpsPass,
648769
RemoveZeroSizedCatArgsPass,
770+
RemovePermuteBeforeMeanPass,
649771
RemovePermutesAroundElementwiseOps,
650772
FuseTransposeOrPermuteOpPairsPass,
651773
RemoveSqueezeViewBeforeElementwiseOps,

backends/cadence/aot/tests/test_remove_ops_passes.py

Lines changed: 195 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
RemoveNopLinalgVectorNormOpPass,
2929
RemoveNopMulOpPass,
3030
RemoveNopSliceOrViewOpPass,
31+
RemovePermuteBeforeMeanPass,
3132
RemovePermutesAroundElementwiseOps,
3233
RemoveSqueezeViewBeforeElementwiseOps,
3334
RemoveToOpsPass,
@@ -1013,3 +1014,197 @@ def test_remove_cat_from_slice_copy_second_input(self) -> None:
10131014

10141015
# Output should remain the same.
10151016
self.assertTrue(torch.equal(graph_module(*inputs)[0], expected_outputs))
1017+
1018+
def test_remove_permute_before_mean_fully_removed(self) -> None:
1019+
"""Permute → relu → mean where non-reduced dims preserve order → fully remove."""
1020+
builder = GraphBuilder()
1021+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1022+
permuted = builder.call_operator(
1023+
op=exir_ops.edge.aten.permute_copy.default, args=(x, [0, 2, 3, 1])
1024+
)
1025+
relu = builder.call_operator(
1026+
op=exir_ops.edge.aten.relu.default, args=(permuted,)
1027+
)
1028+
mean = builder.call_operator(
1029+
op=exir_ops.edge.aten.mean.dim, args=(relu, [1, 2], False)
1030+
)
1031+
builder.output([mean])
1032+
original = builder.get_graph_module()
1033+
gm_before = copy.deepcopy(original)
1034+
1035+
graph_after = cast(
1036+
PassResult, RemovePermuteBeforeMeanPass()(original)
1037+
).graph_module
1038+
1039+
# Permute should be fully removed.
1040+
self.assertEqual(
1041+
count_node(graph_after, exir_ops.edge.aten.permute_copy.default), 0
1042+
)
1043+
1044+
# Mean reduction dims should be remapped to original space.
1045+
mean_nodes = graph_after.graph.find_nodes(
1046+
op="call_function", target=exir_ops.edge.aten.mean.dim
1047+
)
1048+
self.assertEqual(len(mean_nodes), 1)
1049+
self.assertEqual(mean_nodes[0].args[1], [2, 3])
1050+
1051+
validate(
1052+
gm_before,
1053+
graph_after,
1054+
(torch.randn(2, 3, 4, 5),),
1055+
"RemovePermuteBeforeMeanPass",
1056+
)
1057+
1058+
def test_remove_permute_before_mean_sunk_after(self) -> None:
1059+
"""Permute → relu → mean where non-reduced dims reorder → move permute after mean."""
1060+
builder = GraphBuilder()
1061+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1062+
permuted = builder.call_operator(
1063+
op=exir_ops.edge.aten.permute_copy.default, args=(x, [2, 0, 3, 1])
1064+
)
1065+
relu = builder.call_operator(
1066+
op=exir_ops.edge.aten.relu.default, args=(permuted,)
1067+
)
1068+
mean = builder.call_operator(
1069+
op=exir_ops.edge.aten.mean.dim, args=(relu, [2, 3], False)
1070+
)
1071+
builder.output([mean])
1072+
original = builder.get_graph_module()
1073+
gm_before = copy.deepcopy(original)
1074+
1075+
graph_after = cast(
1076+
PassResult, RemovePermuteBeforeMeanPass()(original)
1077+
).graph_module
1078+
1079+
# One permute should remain (after the mean).
1080+
self.assertEqual(
1081+
count_node(graph_after, exir_ops.edge.aten.permute_copy.default), 1
1082+
)
1083+
1084+
# Mean reduction dims should be remapped.
1085+
mean_nodes = graph_after.graph.find_nodes(
1086+
op="call_function", target=exir_ops.edge.aten.mean.dim
1087+
)
1088+
self.assertEqual(len(mean_nodes), 1)
1089+
self.assertEqual(mean_nodes[0].args[1], [3, 1])
1090+
1091+
# The permute should come after the mean, not before.
1092+
permute_nodes = graph_after.graph.find_nodes(
1093+
op="call_function", target=exir_ops.edge.aten.permute_copy.default
1094+
)
1095+
self.assertEqual(permute_nodes[0].args[0], mean_nodes[0])
1096+
self.assertEqual(permute_nodes[0].args[1], [1, 0])
1097+
1098+
validate(
1099+
gm_before,
1100+
graph_after,
1101+
(torch.randn(2, 3, 4, 5),),
1102+
"RemovePermuteBeforeMeanPass",
1103+
)
1104+
1105+
def test_remove_permute_before_mean_keepdim_true(self) -> None:
1106+
"""Permute → relu → mean(keepdim=True) where only reduced dims shuffle → fully remove."""
1107+
builder = GraphBuilder()
1108+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1109+
permuted = builder.call_operator(
1110+
op=exir_ops.edge.aten.permute_copy.default, args=(x, [0, 1, 3, 2])
1111+
)
1112+
relu = builder.call_operator(
1113+
op=exir_ops.edge.aten.relu.default, args=(permuted,)
1114+
)
1115+
mean = builder.call_operator(
1116+
op=exir_ops.edge.aten.mean.dim, args=(relu, [2, 3], True)
1117+
)
1118+
builder.output([mean])
1119+
original = builder.get_graph_module()
1120+
gm_before = copy.deepcopy(original)
1121+
1122+
graph_after = cast(
1123+
PassResult, RemovePermuteBeforeMeanPass()(original)
1124+
).graph_module
1125+
1126+
# Permute fully removed (only reduced dims were shuffled).
1127+
self.assertEqual(
1128+
count_node(graph_after, exir_ops.edge.aten.permute_copy.default), 0
1129+
)
1130+
1131+
validate(
1132+
gm_before,
1133+
graph_after,
1134+
(torch.randn(2, 3, 4, 5),),
1135+
"RemovePermuteBeforeMeanPass",
1136+
)
1137+
1138+
def test_remove_permute_before_mean_keepdim_true_sunk(self) -> None:
1139+
"""Permute → relu → mean(keepdim=True) where non-reduced dims move → sink permute."""
1140+
builder = GraphBuilder()
1141+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1142+
permuted = builder.call_operator(
1143+
op=exir_ops.edge.aten.permute_copy.default, args=(x, [0, 2, 3, 1])
1144+
)
1145+
relu = builder.call_operator(
1146+
op=exir_ops.edge.aten.relu.default, args=(permuted,)
1147+
)
1148+
mean = builder.call_operator(
1149+
op=exir_ops.edge.aten.mean.dim, args=(relu, [1, 2], True)
1150+
)
1151+
builder.output([mean])
1152+
original = builder.get_graph_module()
1153+
gm_before = copy.deepcopy(original)
1154+
1155+
graph_after = cast(
1156+
PassResult, RemovePermuteBeforeMeanPass()(original)
1157+
).graph_module
1158+
1159+
# One permute should remain (sunk after mean).
1160+
self.assertEqual(
1161+
count_node(graph_after, exir_ops.edge.aten.permute_copy.default), 1
1162+
)
1163+
1164+
# The post-mean permute uses the original perm since keepdim=True.
1165+
permute_nodes = graph_after.graph.find_nodes(
1166+
op="call_function", target=exir_ops.edge.aten.permute_copy.default
1167+
)
1168+
self.assertEqual(permute_nodes[0].args[1], [0, 2, 3, 1])
1169+
1170+
validate(
1171+
gm_before,
1172+
graph_after,
1173+
(torch.randn(2, 3, 4, 5),),
1174+
"RemovePermuteBeforeMeanPass",
1175+
)
1176+
1177+
def test_remove_permute_before_mean_no_permute(self) -> None:
1178+
"""No permute before mean → no change."""
1179+
builder = GraphBuilder()
1180+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1181+
relu = builder.call_operator(op=exir_ops.edge.aten.relu.default, args=(x,))
1182+
mean = builder.call_operator(
1183+
op=exir_ops.edge.aten.mean.dim, args=(relu, [2, 3], False)
1184+
)
1185+
builder.output([mean])
1186+
original = builder.get_graph_module()
1187+
1188+
result = cast(PassResult, RemovePermuteBeforeMeanPass()(original))
1189+
self.assertFalse(result.modified)
1190+
1191+
def test_remove_permute_before_mean_multi_user(self) -> None:
1192+
"""Permute with multiple users → no change."""
1193+
builder = GraphBuilder()
1194+
x = builder.placeholder("x", torch.randn(2, 3, 4, 5))
1195+
permuted = builder.call_operator(
1196+
op=exir_ops.edge.aten.permute_copy.default, args=(x, [0, 2, 3, 1])
1197+
)
1198+
relu = builder.call_operator(
1199+
op=exir_ops.edge.aten.relu.default, args=(permuted,)
1200+
)
1201+
mean = builder.call_operator(
1202+
op=exir_ops.edge.aten.mean.dim, args=(relu, [1, 2], False)
1203+
)
1204+
# Second user of the permute prevents optimization.
1205+
neg = builder.call_operator(op=exir_ops.edge.aten.neg.default, args=(permuted,))
1206+
builder.output([mean, neg])
1207+
original = builder.get_graph_module()
1208+
1209+
result = cast(PassResult, RemovePermuteBeforeMeanPass()(original))
1210+
self.assertFalse(result.modified)

0 commit comments

Comments
 (0)