@@ -1358,6 +1358,73 @@ def test_unsupported_parent_not_swapped(self) -> None:
13581358
13591359 self .assertFalse (result .modified )
13601360
1361+ def test_swap_additional_unary_target (self ) -> None :
1362+ x_data = torch .randn (4 , 60 , 1 , 1 )
1363+ builder = GraphBuilder ()
1364+ x = builder .placeholder ("x" , x_data )
1365+ relu = builder .call_operator (exir_ops .edge .aten .relu .default , args = (x ,))
1366+ sliced = builder .call_operator (
1367+ exir_ops .edge .aten .slice_copy .Tensor ,
1368+ args = (relu , 0 , 0 , 4 , 2 ),
1369+ )
1370+ builder .output ([sliced ])
1371+ gm = builder .get_graph_module ()
1372+
1373+ result = transform_and_check_numerics (
1374+ gm ,
1375+ (x_data ,),
1376+ PropagateSlice (additional_unary_targets = [exir_ops .edge .aten .relu .default ]),
1377+ )
1378+
1379+ self .assertTrue (result .modified )
1380+ slice_nodes = gm .graph .find_nodes (
1381+ op = "call_function" , target = exir_ops .edge .aten .slice_copy .Tensor
1382+ )
1383+ self .assertEqual (len (slice_nodes ), 1 )
1384+ relu_nodes = gm .graph .find_nodes (
1385+ op = "call_function" , target = exir_ops .edge .aten .relu .default
1386+ )
1387+ self .assertEqual (len (relu_nodes ), 1 )
1388+ self .assertIs (relu_nodes [0 ].args [0 ], slice_nodes [0 ])
1389+ self .assertEqual (list (relu_nodes [0 ].meta ["val" ].shape ), [2 , 60 , 1 , 1 ])
1390+
1391+ def test_swap_additional_binary_target (self ) -> None :
1392+ lhs_data = torch .randn (1 , 60 , 1 , 1 )
1393+ rhs_data = torch .randn (4 , 60 , 1 , 1 )
1394+ builder = GraphBuilder ()
1395+ lhs = builder .placeholder ("lhs" , lhs_data )
1396+ rhs = builder .placeholder ("rhs" , rhs_data )
1397+ sub = builder .call_operator (
1398+ exir_ops .edge .aten .sub .Tensor ,
1399+ args = (lhs , rhs ),
1400+ )
1401+ sliced = builder .call_operator (
1402+ exir_ops .edge .aten .slice_copy .Tensor ,
1403+ args = (sub , 0 , 0 , 4 , 2 ),
1404+ )
1405+ builder .output ([sliced ])
1406+ gm = builder .get_graph_module ()
1407+
1408+ result = transform_and_check_numerics (
1409+ gm ,
1410+ (lhs_data , rhs_data ),
1411+ PropagateSlice (additional_binary_targets = [exir_ops .edge .aten .sub .Tensor ]),
1412+ )
1413+
1414+ self .assertTrue (result .modified )
1415+ slice_nodes = gm .graph .find_nodes (
1416+ op = "call_function" , target = exir_ops .edge .aten .slice_copy .Tensor
1417+ )
1418+ self .assertEqual (len (slice_nodes ), 1 )
1419+ self .assertEqual (slice_nodes [0 ].args [0 ].name , "rhs" )
1420+ sub_nodes = gm .graph .find_nodes (
1421+ op = "call_function" , target = exir_ops .edge .aten .sub .Tensor
1422+ )
1423+ self .assertEqual (len (sub_nodes ), 1 )
1424+ self .assertIs (sub_nodes [0 ].args [0 ], lhs .node )
1425+ self .assertIs (sub_nodes [0 ].args [1 ], slice_nodes [0 ])
1426+ self .assertEqual (list (sub_nodes [0 ].meta ["val" ].shape ), [2 , 60 , 1 , 1 ])
1427+
13611428 def test_swap_broadcast_mul_slice_on_broadcast_dim (self ) -> None :
13621429 """[1,60,1,1] * [4,1,1,1] → [4,60,1,1] → slice(dim=0, step=2)
13631430 Only the [4,1,1,1] input should be sliced."""
0 commit comments