@@ -1034,39 +1034,6 @@ def test_permute_unsqueeze_copy_neg_dim_mul_squeeze_copy_permute(self) -> None:
10341034 "permute_unsqueeze_copy_neg_dim_mul_squeeze_copy_permute" ,
10351035 )
10361036
1037- def test_unsqueeze_at_moved_position (self ) -> None :
1038- """The permutation moves the unsqueeze position (P[index] != index), so
1039- the adapted permutation must be built from P[index], not index."""
1040- builder = GraphBuilder ()
1041- x_data = torch .randn (1 , 8 , 16 )
1042- x = builder .placeholder ("x" , x_data )
1043- p1 = builder .call_operator (
1044- op = exir_ops .edge .aten .permute_copy .default , args = (x , [0 , 2 , 1 ])
1045- )
1046- u = builder .call_operator (
1047- op = exir_ops .edge .aten .unsqueeze_copy .default , args = (p1 , 1 )
1048- )
1049- mul = builder .call_operator (op = exir_ops .edge .aten .mul .Tensor , args = (u , u ))
1050- sq = builder .call_operator (
1051- op = exir_ops .edge .aten .squeeze_copy .dim , args = (mul , 1 )
1052- )
1053- p2 = builder .call_operator (
1054- op = exir_ops .edge .aten .permute_copy .default , args = (sq , [0 , 2 , 1 ])
1055- )
1056- builder .output ([p2 ])
1057- original = builder .get_graph_module ()
1058- gm_before = copy .deepcopy (original )
1059-
1060- p = RemovePermutesAroundElementwiseOps ()
1061- result = cast (PassResult , p (original ))
1062- self .assertTrue (result .modified )
1063- self .assertEqual (
1064- count_node (result .graph_module , exir_ops .edge .aten .permute_copy .default ), 0
1065- )
1066- validate_numerics (
1067- gm_before , result .graph_module , [x_data ], "UnsqueezeAtMovedPosition"
1068- )
1069-
10701037 def test_upstream_view_rank_mismatch_no_crash (self ) -> None :
10711038 """Regression test for IndexError when a squeeze/unsqueeze view_copy
10721039 is reached via upstream traversal with a permutation whose rank does
@@ -1170,6 +1137,124 @@ def test_permutation_sink_view_splitting_the_non_unit_dim(self) -> None:
11701137 "permutation_sink_view_splitting_the_non_unit_dim" ,
11711138 )
11721139
1140+ def test_upstream_squeeze_view_rank_mismatch_no_crash (self ) -> None :
1141+ """Regression test for IndexError when a squeeze view_copy is reached
1142+ via upstream traversal with a permutation at the view's output rank.
1143+
1144+ Graph:
1145+ y [8, 4, 1] x [4, 8]
1146+ | |
1147+ view_copy (squeeze 3D→2D) permute [1, 0]
1148+ [8, 4] [8, 4]
1149+ \\ /
1150+ ---- add (2D) -----------
1151+ |
1152+ permute [1, 0]
1153+ |
1154+ output
1155+
1156+ `visit` reaches the view_copy from `add` with the 2D permutation
1157+ [1, 0], but the view's input is 3D, so the squeezed position (2) is
1158+ out of range for the permutation. Before the fix,
1159+ _adapt_permute_across_view crashed with IndexError: list index out
1160+ of range."""
1161+ builder = GraphBuilder ()
1162+ x_data = torch .randn (4 , 8 )
1163+ y_data = torch .randn (8 , 4 , 1 )
1164+ x = builder .placeholder ("x" , x_data )
1165+ y = builder .placeholder ("y" , y_data )
1166+ # Squeeze via view_copy: [8, 4, 1] → [8, 4]
1167+ view_sq = builder .call_operator (
1168+ op = exir_ops .edge .aten .view_copy .default , args = (y , [8 , 4 ])
1169+ )
1170+ # Start permute: [4, 8] → [8, 4]
1171+ p1 = builder .call_operator (
1172+ op = exir_ops .edge .aten .permute_copy .default , args = (x , [1 , 0 ])
1173+ )
1174+ add = builder .call_operator (
1175+ op = exir_ops .edge .aten .add .Tensor , args = (p1 , view_sq )
1176+ )
1177+ # End permute: [8, 4] → [4, 8]
1178+ p2 = builder .call_operator (
1179+ op = exir_ops .edge .aten .permute_copy .default , args = (add , [1 , 0 ])
1180+ )
1181+ builder .output ([p2 ])
1182+ original = builder .get_graph_module ()
1183+ gm_before = copy .deepcopy (original )
1184+
1185+ # Should not crash, and should skip the subgraph due to rank mismatch
1186+ p = RemovePermutesAroundElementwiseOps ()
1187+ result = cast (PassResult , p (original ))
1188+ self .assertFalse (result .modified )
1189+ self .assertEqual (
1190+ count_node (result .graph_module , exir_ops .edge .aten .permute_copy .default ), 2
1191+ )
1192+ validate_numerics (
1193+ gm_before ,
1194+ result .graph_module ,
1195+ [x_data , y_data ],
1196+ "upstream_squeeze_view_rank_mismatch_no_crash" ,
1197+ )
1198+
1199+ def test_broadcast_rank_increase_no_crash (self ) -> None :
1200+ """Regression test for IndexError when broadcasting raises a node's
1201+ rank above the rank of the permutation it is visited with.
1202+
1203+ Graph:
1204+ x [1, 8] full([1, 1, 1])
1205+ | |
1206+ permute [1, 0] |
1207+ [8, 1] |
1208+ \\ /
1209+ ------ add ------------
1210+ [1, 8, 1] (rank 3)
1211+ |
1212+ slice_copy(dim=2)
1213+ |
1214+ view_copy -> [8] (permutation sink)
1215+
1216+ `add` is visited with the rank-2 permutation [1, 0] but broadcasts to
1217+ rank 3. The numel-1 `full` input and the sink `view_copy` both let
1218+ traversal terminate without ever meeting a rank-matched permute, so the
1219+ subgraph closed with a rank-2 permutation on rank-3 nodes. Before the
1220+ fix, update_slice_copy did `start_permute[2]` and raised
1221+ IndexError: list index out of range."""
1222+ builder = GraphBuilder ()
1223+ x_data = torch .randn (1 , 8 )
1224+ x = builder .placeholder ("x" , x_data )
1225+ p1 = builder .call_operator (
1226+ op = exir_ops .edge .aten .permute_copy .default , args = (x , [1 , 0 ])
1227+ )
1228+ ones = builder .call_operator (
1229+ op = exir_ops .edge .aten .full .default , args = ([1 , 1 , 1 ], 1.0 )
1230+ )
1231+ # Broadcast [8, 1] + [1, 1, 1] -> [1, 8, 1]: rank 3 under a rank 2 permute
1232+ add = builder .call_operator (op = exir_ops .edge .aten .add .Tensor , args = (p1 , ones ))
1233+ sl = builder .call_operator (
1234+ op = exir_ops .edge .aten .slice_copy .Tensor , args = (add , 2 , 0 , 1 )
1235+ )
1236+ # Sink view: input [1, 8, 1] has a single non-unit dim
1237+ sink = builder .call_operator (
1238+ op = exir_ops .edge .aten .view_copy .default , args = (sl , [8 ])
1239+ )
1240+ builder .output ([sink ])
1241+ original = builder .get_graph_module ()
1242+ gm_before = copy .deepcopy (original )
1243+
1244+ # Should not crash, and should skip the subgraph due to rank mismatch
1245+ p = RemovePermutesAroundElementwiseOps ()
1246+ result = cast (PassResult , p (original ))
1247+ self .assertFalse (result .modified )
1248+ self .assertEqual (
1249+ count_node (result .graph_module , exir_ops .edge .aten .permute_copy .default ), 1
1250+ )
1251+ validate_numerics (
1252+ gm_before ,
1253+ result .graph_module ,
1254+ [x_data ],
1255+ "broadcast_rank_increase_no_crash" ,
1256+ )
1257+
11731258
11741259# ──────────────────────────────────────────────────────────────────────
11751260# Tests for RemovePermutesAroundElementwiseOps
0 commit comments