@@ -1034,6 +1034,39 @@ 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+
10371070 def test_upstream_view_rank_mismatch_no_crash (self ) -> None :
10381071 """Regression test for IndexError when a squeeze/unsqueeze view_copy
10391072 is reached via upstream traversal with a permutation whose rank does
0 commit comments