|
28 | 28 | RemoveNopLinalgVectorNormOpPass, |
29 | 29 | RemoveNopMulOpPass, |
30 | 30 | RemoveNopSliceOrViewOpPass, |
| 31 | + RemovePermuteBeforeMeanPass, |
31 | 32 | RemovePermutesAroundElementwiseOps, |
32 | 33 | RemoveSqueezeViewBeforeElementwiseOps, |
33 | 34 | RemoveToOpsPass, |
@@ -1013,3 +1014,197 @@ def test_remove_cat_from_slice_copy_second_input(self) -> None: |
1013 | 1014 |
|
1014 | 1015 | # Output should remain the same. |
1015 | 1016 | 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